{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":6259291,"sourceType":"datasetVersion","datasetId":3597559},{"sourceId":7392733,"sourceType":"datasetVersion","datasetId":4297749},{"sourceId":7392775,"sourceType":"datasetVersion","datasetId":4297782},{"sourceId":7402356,"sourceType":"datasetVersion","datasetId":4304475},{"sourceId":7403069,"sourceType":"datasetVersion","datasetId":4304949},{"sourceId":7447509,"sourceType":"datasetVersion","datasetId":4334995},{"sourceId":7450712,"sourceType":"datasetVersion","datasetId":4336944},{"sourceId":7465251,"sourceType":"datasetVersion","datasetId":4317718},{"sourceId":7581697,"sourceType":"datasetVersion","datasetId":4413439},{"sourceId":7581715,"sourceType":"datasetVersion","datasetId":4413451},{"sourceId":7581720,"sourceType":"datasetVersion","datasetId":4413454},{"sourceId":7617703,"sourceType":"datasetVersion","datasetId":4436702},{"sourceId":7866794,"sourceType":"datasetVersion","datasetId":4615431},{"sourceId":7867343,"sourceType":"datasetVersion","datasetId":4615834},{"sourceId":7867344,"sourceType":"datasetVersion","datasetId":4615835},{"sourceId":7867345,"sourceType":"datasetVersion","datasetId":4615836},{"sourceId":7867346,"sourceType":"datasetVersion","datasetId":4615837},{"sourceId":7867347,"sourceType":"datasetVersion","datasetId":4615838},{"sourceId":7955509,"sourceType":"datasetVersion","datasetId":4679124},{"sourceId":7955510,"sourceType":"datasetVersion","datasetId":4679125},{"sourceId":7979206,"sourceType":"datasetVersion","datasetId":4596419},{"sourceId":7979733,"sourceType":"datasetVersion","datasetId":4696444},{"sourceId":7979752,"sourceType":"datasetVersion","datasetId":4696459},{"sourceId":7979770,"sourceType":"datasetVersion","datasetId":4696472},{"sourceId":7979791,"sourceType":"datasetVersion","datasetId":4696485},{"sourceId":7979920,"sourceType":"datasetVersion","datasetId":4696576},{"sourceId":7979949,"sourceType":"datasetVersion","datasetId":4696595},{"sourceId":7981321,"sourceType":"datasetVersion","datasetId":4419986},{"sourceId":7987646,"sourceType":"datasetVersion","datasetId":4701635},{"sourceId":7989915,"sourceType":"datasetVersion","datasetId":4703503},{"sourceId":8015225,"sourceType":"datasetVersion","datasetId":4435105},{"sourceId":8040912,"sourceType":"datasetVersion","datasetId":4652811},{"sourceId":158958765,"sourceType":"kernelVersion"},{"sourceId":159333316,"sourceType":"kernelVersion"},{"sourceId":159396114,"sourceType":"kernelVersion"},{"sourceId":162276846,"sourceType":"kernelVersion"},{"sourceId":162740849,"sourceType":"kernelVersion"},{"sourceId":168960092,"sourceType":"kernelVersion"},{"sourceId":170023803,"sourceType":"kernelVersion"},{"sourceId":170023919,"sourceType":"kernelVersion"},{"sourceId":170024030,"sourceType":"kernelVersion"},{"sourceId":170091640,"sourceType":"kernelVersion"},{"sourceId":170092156,"sourceType":"kernelVersion"},{"sourceId":170476992,"sourceType":"kernelVersion"},{"sourceId":170485175,"sourceType":"kernelVersion"},{"sourceId":170565780,"sourceType":"kernelVersion"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport warnings\nwarnings.simplefilter('ignore')\npd.set_option('display.max_columns', 100)\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nfrom scipy.special import softmax\nimport multiprocessing\nimport copy\n\nimport matplotlib.pyplot as plt\nfrom glob import glob\nfrom tqdm import tqdm\nimport random\nimport os\ndef seed_everything(seed=2024):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\nimport time\ndebug = True\nseed_everything()\nDEVICE = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:26:24.87485Z","iopub.execute_input":"2024-04-06T07:26:24.875556Z","iopub.status.idle":"2024-04-06T07:26:29.931483Z","shell.execute_reply.started":"2024-04-06T07:26:24.875524Z","shell.execute_reply":"2024-04-06T07:26:29.930658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/kaggle-kl-div')\nfrom kaggle_kl_div import score as kaggle_metric","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:26:29.933046Z","iopub.execute_input":"2024-04-06T07:26:29.933442Z","iopub.status.idle":"2024-04-06T07:26:29.955212Z","shell.execute_reply.started":"2024-04-06T07:26:29.933417Z","shell.execute_reply":"2024-04-06T07:26:29.954519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%capture\n!ln -s /kaggle/input/timm092/pytorch-image-models-0.9.2/timm\nimport timm\n# !pip install pfio\n!pip install einops\n!pip install albumentations==1.4.3\n","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:26:29.956197Z","iopub.execute_input":"2024-04-06T07:26:29.956522Z","iopub.status.idle":"2024-04-06T07:27:05.690212Z","shell.execute_reply.started":"2024-04-06T07:26:29.956492Z","shell.execute_reply":"2024-04-06T07:27:05.68881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append(\"/kaggle/input/kaggle-kl-div\")\nfrom kaggle_kl_div import score as kaggle_metric\nimport pandas as pd\nimport numpy as np\nfrom scipy.special import softmax\nfrom scipy.optimize import minimize\nfrom glob import glob\nfrom tqdm import tqdm\nimport warnings\nwarnings.simplefilter('ignore')\ndef sigmoid(x):\n    return 1/(1 + np.exp(-x))\npd.set_option('display.max_columns', 100)\nfrom sklearn.metrics import roc_auc_score, confusion_matrix, mean_squared_error, average_precision_score, accuracy_score\n\ndef calc_metric(oof):\n    true = oof[true_cols]\n    pred = oof[pred_cols]\n    pred[pred_cols] = softmax(pred[pred_cols].values, 1)\n    true['id'] = list(range(len(true)))\n    pred['id'] = list(range(len(pred)))\n    true.columns = [0,1,2,3,4,5,'id']\n    pred.columns = [0,1,2,3,4,5,'id']\n    score = kaggle_metric(solution=true, submission=pred, row_id_column_name='id')\n    return score\n\ndef kl_divergence_for_scipy(final_preds):\n    epsilon = 1e-10\n    final_preds = softmax(final_preds, 1) + epsilon\n    scores = []\n    for i in range(len(pred_cols)):\n        pr = final_preds[:, i]\n        score = np.mean(true[:, i] * np.log(true[:, i] / pr))\n        scores.append(score)\n    return np.sum(scores)\n\ndef kl_divergence(true, final_preds):\n    epsilon = 1e-10\n    final_preds = softmax(final_preds, 1) + epsilon\n    scores = []\n    for i in range(len(pred_cols)):\n        pr = final_preds[:, i]\n        score = np.mean(true[:, i] * np.log(true[:, i] / pr))\n        scores.append(score)\n    return np.sum(scores)\n\ndef loss_fn(weights):\n    final_preds = 0\n    weights = np.array(weights)/np.sum(weights)\n    for weight, pred in zip(weights, preds):\n        final_preds += weight*pred\n    score = kl_divergence_for_scipy(final_preds)\n    return score\n    \ndef round_to_nearest(x, base=0.05):\n    return base * np.round(x/base)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:27:05.693704Z","iopub.execute_input":"2024-04-06T07:27:05.694115Z","iopub.status.idle":"2024-04-06T07:27:06.435434Z","shell.execute_reply.started":"2024-04-06T07:27:05.694072Z","shell.execute_reply":"2024-04-06T07:27:06.434369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%cd /kaggle/working","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:27:06.436819Z","iopub.execute_input":"2024-04-06T07:27:06.437553Z","iopub.status.idle":"2024-04-06T07:27:06.44606Z","shell.execute_reply.started":"2024-04-06T07:27:06.437526Z","shell.execute_reply":"2024-04-06T07:27:06.445036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cols = ['seizure_vote','lpd_vote','gpd_vote','lrda_vote','grda_vote','other_vote']\n","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:27:06.447218Z","iopub.execute_input":"2024-04-06T07:27:06.447486Z","iopub.status.idle":"2024-04-06T07:27:06.459551Z","shell.execute_reply.started":"2024-04-06T07:27:06.447463Z","shell.execute_reply":"2024-04-06T07:27:06.458634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cp -r /kaggle/input/hms-pipeline ./","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:27:06.460971Z","iopub.execute_input":"2024-04-06T07:27:06.461342Z","iopub.status.idle":"2024-04-06T07:27:08.863854Z","shell.execute_reply.started":"2024-04-06T07:27:06.46131Z","shell.execute_reply":"2024-04-06T07:27:08.862598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir -p /kaggle/working/hms/results","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:27:08.865265Z","iopub.execute_input":"2024-04-06T07:27:08.86556Z","iopub.status.idle":"2024-04-06T07:27:09.820755Z","shell.execute_reply.started":"2024-04-06T07:27:08.865534Z","shell.execute_reply":"2024-04-06T07:27:09.819675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile /kaggle/working/hms-pipeline/hms_pipeline/scripts/train_one_fold.py\n\nimport os\nimport shutil\nfrom multiprocessing import cpu_count\nimport sys\nimport datetime\nimport time\nsys.path.append('../')\nsys.path.append('/home/acc12347av/ml_pipeline')\n# sys.path.insert(0, \"src/pytorch-lightning\")\n\nfrom pytorch_lightning import Trainer, seed_everything\n# from pytorch_lightning.plugins import DDPPlugin\nfrom pytorch_lightning.strategies import DDPStrategy\nfrom pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping, LearningRateMonitor\nfrom pytorch_lightning.loggers import WandbLogger\nfrom pytorch_lightning.loggers.csv_logs import CSVLogger\n\nfrom pdb import set_trace as st\nimport warnings\nwarnings.simplefilter('ignore')\n\nimport argparse\ndef parse_args():\n    parser = argparse.ArgumentParser()\n    parser.add_argument(\"--config\", '-c', type=str, default='Test', help=\"config name\")\n    parser.add_argument(\"--type\", '-t', type=str, default='classification')\n    parser.add_argument(\"--gpu\", '-g', type=str, default='nochange')\n    parser.add_argument(\"--debug\", action='store_true', help=\"debug\")\n    parser.add_argument(\"--wandb\", action='store_true', help=\"wandb\")\n    parser.add_argument(\"--fold\", '-f', type=int, default=0, help=\"fold num\")\n    parser.add_argument(\"--epochs\", '-e', type=int, default=1, help=\"num epochs\")\n    parser.add_argument(\"--use_row\", type=int, default=2, help=\"google spread sheet row\")\n    return parser.parse_args()\n\ndef load_model(model, pretrained_path, num_classes=1, skip_attn=False):\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    checkpoint = torch.load(pretrained_path, map_location=device)\n    print('load pretrained model from', pretrained_path)\n    if skip_attn:\n        checkpoint['model'] = {k: v for k, v in checkpoint['model'].items() if 'attn_mask' not in str(k)}\n\n    if isinstance(model, torch.nn.DataParallel):\n        model.module.load_state_dict(checkpoint['model'])\n    else:\n        model.load_state_dict(checkpoint['model'])\n\n    if 'SwinTransformer' in str(type(model)):\n        model.head = nn.Linear(in_features=1536, out_features=num_classes, bias=True)\n        print('change num_classes to', num_classes)\n    elif 'Cait' in str(type(model)):\n        model.head = nn.Linear(in_features=768, out_features=num_classes, bias=True)\n        print('change num_classes to', num_classes)\n\n\nCONFIG_COL_N = 2\nTIME_COL_N = 7\n\nif __name__ == \"__main__\":\n    start = time.time()\n    args = parse_args()\n    if args.type == 'classification':\n        from src.configs import *\n    elif args.type == 'seg':\n        from src.seg_configs import *\n    elif args.type == 'effdet':\n        from src.effdet_configs import *\n    elif args.type == 'nlp':\n        from src.nlp_configs import *\n    elif args.type == 'mlp_with_nlp':\n        from src.mlp_with_nlp_configs import *\n    elif args.type == 'mlp':\n        from src.mlp_configs import *\n    elif args.type == 'gnn':\n        from src.gnn_configs import *\n\n    try:\n        cfg = eval(args.config)(args.fold)\n    except Exception as e:\n        # print('eval(args.config)(args.fold): error')\n        # print(e)\n        cfg = eval(args.config)()\n    use_wandb = False\n\n    cfg.gpu = 'small'\n    cfg.epochs = args.epochs\n    if cfg.gpu == 'big':\n        devices = 8\n    elif cfg.gpu == 'small':\n        devices = 1\n    elif cfg.gpu == 'v100':\n        devices = 4\n    else:\n        raise\n    if cfg.gpu not in ['big', 'small']: # A100じゃなかったら(= V100だったら)\n        cfg.batch_size = cfg.batch_size // 2\n        cfg.lr /= 2\n        if cfg.batch_size < 16:\n            cfg.grad_accumulations *= 2\n            cfg.lr *= 2\n\n    if cfg.inference_only:\n        exit()\n        raise\n    if cfg.train_by_all_data & (args.fold != 0):\n        exit()\n    cfg.fold = args.fold\n    if cfg.seed is None:\n        now = datetime.datetime.now()\n        cfg.seed = int(now.strftime('%s'))\n\n    if use_wandb:\n        import wandb\n        wandb.login()\n    else:\n        os.environ['WANDB_MODE'] = 'offline'\n\n    RESULTS_PATH_BASE = f'/kaggle/working/hms/results'\n    wandb_save_dir = '/kaggle/working'\n\n    if args.type == 'classification':\n        from src.lightning.lightning_modules.classification import MyLightningModule\n        from src.lightning.data_modules.classification import MyDataModule\n    elif args.type == 'seg':\n        from src.lightning.lightning_modules.segmentation import MyLightningModule\n        from src.lightning.data_modules.segmentation import MyDataModule\n    elif args.type == 'effdet':\n        from src.lightning.lightning_modules.effdet import MyLightningModule\n        from src.lightning.data_modules.effdet import MyDataModule\n    elif args.type == 'nlp':\n        from src.lightning.lightning_modules.nlp import MyLightningModule\n        from src.lightning.data_modules.nlp import MyDataModule\n    elif args.type == 'mlp_with_nlp':\n        from src.lightning.lightning_modules.mlp_with_nlp import MyLightningModule\n        from src.lightning.data_modules.mlp_with_nlp import MyDataModule\n    elif args.type == 'mlp':\n        from src.lightning.lightning_modules.mlp import MyLightningModule\n        from src.lightning.data_modules.mlp import MyDataModule\n    elif args.type == 'gnn':\n        from src.lightning.lightning_modules.gnn import MyLightningModule\n        from src.lightning.data_modules.gnn import MyDataModule\n\n    if args.debug:\n        cfg.epochs = 1\n        cfg.n_cpu = 1\n        n_gpu = 1\n    else:\n        n_gpu = torch.cuda.device_count()\n        cfg.n_cpu = n_gpu * np.min([cpu_count(), cfg.batch_size])\n\n    print(f'\\n----------------------- Config -----------------------')\n    from datetime import datetime\n    config_str = f'time: {datetime.now().strftime(\"%m-%d %H:%M\")}, '\n    for k, v in vars(cfg).items():\n        if (k == 'model') | (k == 'teacher_models') | (('df' in k) & ('path' not in k)):\n            continue\n        if (k == 'label_features') and (len(v) > 100):\n            print(f'\\t{k} len: {len(v)}')\n            config_str += f'{k} len: {len(v)}, '\n        else:\n            if type(v) != int: v = str(v).replace(\"\\n\", \"\")\n            print(f'\\t{k}: {v}')\n            config_str += f'{k}: {v}, '\n    config_str += f'train_df len: {len(cfg.train_df)}'\n    if cfg.valid_df is not None:\n        config_str += f'valid_df len: {len(cfg.valid_df)}'\n\n    print(f'----------------------- Config -----------------------\\n')\n    print('config_str:', config_str)\n\n    if args.type not in ['nlp', 'mlp_with_nlp', 'mlp', 'gnn']:\n        if type(cfg.image_size) == int:\n            cfg.image_size = (cfg.image_size, cfg.image_size)\n        cfg.transform = cfg.transform(cfg.image_size)\n    if args.type == 'effdet':\n        from src.utils.augmentations.det_tta_augmentation import *\n        if cfg.tta_transforms == '':\n            cfg.tta_transforms = [TTACompose([])]\n        else:\n            tta_transforms = []\n            for tta_transform in cfg.tta_transforms.split(','):\n                if tta_transform == 'hflip':\n                    tta_class = TTAHorizontalFlip\n                elif tta_transform == 'hflip':\n                    tta_class = TTAVerticalFlip\n                elif tta_transform == 'rotate90':\n                    tta_class = TTARotate90\n                else:\n                    raise\n                tta_transforms.append([tta_class(size), None])\n            tta_transforms = product(tta_transforms)\n            cfg.tta_transforms = []\n            for tta_combination in tta_transforms:\n                cfg.tta_transforms.append(TTACompose([tta_transform for tta_transform in tta_combination if tta_transform]))\n\n    if cfg.pretrained_path is not None:\n        torch_state_dict = torch.load(cfg.pretrained_path, map_location=torch.device('cpu'))\n        cfg.model.load_state_dict(torch_state_dict)\n        print('pretrain_path:', cfg.pretrained_path)\n\n\n    seed_everything(cfg.seed)\n    OUTPUT_PATH = f'{RESULTS_PATH_BASE}/{args.config}'\n    cfg.output_path = OUTPUT_PATH\n    logger = CSVLogger(save_dir=OUTPUT_PATH, name=f\"fold_{args.fold}\")\n    wandb_logger = WandbLogger(name=f'{args.config}_fold{args.fold}', project=cfg.compe, group=args.config, offline=not use_wandb, save_dir=wandb_save_dir)\n\n    monitor = 'val_metric'\n    checkpoint_callback = ModelCheckpoint(\n        dirpath=OUTPUT_PATH, filename=f\"fold_{args.fold}\", auto_insert_metric_name=False,\n        save_top_k=cfg.save_top_k, monitor=monitor, mode='max', verbose=True, save_last=False)\n    # checkpoint_callback.CHECKPOINT_NAME_LAST = f'last_fold{args.fold}.ckpt'\n    # checkpoint_callback.CHECKPOINT_NAME_LAST = \"last_{epoch}\"+f'_fold{args.fold}'\n    # checkpoint_callback.FILE_EXTENSION = \"_epoch_{epoch}\"+f'.ckpt'\n\n    early_stop_callback = EarlyStopping(patience=cfg.early_stop_patience,\n        monitor=monitor, mode='max', verbose=True)\n    lr_monitor = LearningRateMonitor(logging_interval='epoch')\n    strategy = 'ddp'\n    # https://pytorch-lightning.readthedocs.io/en/latest/common/trainer.html\n    trainer = Trainer(\n        # auto_lr_find=True,\n        # auto_scale_batch_size=\"binsearch\",\n        max_epochs=cfg.epochs,\n        accumulate_grad_batches=cfg.grad_accumulations,\n        precision=16 if cfg.fp16 else 32,\n        # amp_backend='native',\n        deterministic=False,\n        benchmark=True,\n        limit_train_batches=1.0,\n        limit_val_batches=1.0,\n        callbacks=[checkpoint_callback, early_stop_callback, lr_monitor],\n        # logger=[logger],\n        logger=[logger],\n        # plugins=DDPPlugin(find_unused_parameters=(args.type=='nlp')|(args.type=='seg')),\n        # plugins=DDPStrategy(find_unused_parameters=(args.type=='nlp')|(args.type=='seg')),\n        # plugins=DDPPlugin(find_unused_parameters=True),\n        # sync_batchnorm=cfg.sync_batchnorm,\n        enable_progress_bar=False,\n        # resume_from_checkpoint=f'{OUTPUT_PATH}/fold_{args.fold}.ckpt' if cfg.resume else None,\n        accelerator='gpu',\n        # strategy=strategy,\n        devices=devices,\n        reload_dataloaders_every_n_epochs=getattr(cfg, 'reload_dataloaders_every_n_epochs', 0),\n        # gradient_clip_val=0.5, gradient_clip_algorithm=\"value\"\n        fast_dev_run=args.debug,\n    )\n    model = MyLightningModule(cfg)\n    datamodule = MyDataModule(cfg)\n\n    print('start training.')\n    trainer.fit(model, datamodule=datamodule)\n    os.system(f'ls {OUTPUT_PATH}/')\n    torch.save(model.model.state_dict(), f'{OUTPUT_PATH}/last_fold{args.fold}.ckpt')\n#     best_model_path = checkpoint_callback.best_model_path\n#     best_model = model.load_from_checkpoint(cfg=cfg, checkpoint_path=best_model_path)\n#     torch.save(best_model.model.state_dict(), f'{OUTPUT_PATH}/fold_{args.fold}.ckpt')\n    os.system(f'rm {OUTPUT_PATH}/fold_{args.fold}.ckpt')\n","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:27:09.822469Z","iopub.execute_input":"2024-04-06T07:27:09.822801Z","iopub.status.idle":"2024-04-06T07:27:09.837402Z","shell.execute_reply.started":"2024-04-06T07:27:09.822774Z","shell.execute_reply":"2024-04-06T07:27:09.836517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%writefile /kaggle/working/hms-pipeline/hms_pipeline/src/configs.py\n","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:27:09.841632Z","iopub.execute_input":"2024-04-06T07:27:09.841957Z","iopub.status.idle":"2024-04-06T07:27:09.850604Z","shell.execute_reply.started":"2024-04-06T07:27:09.841934Z","shell.execute_reply":"2024-04-06T07:27:09.849757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile /kaggle/working/hms-pipeline/hms_pipeline/src/configs.py\n\nfrom pathlib import Path\nfrom pprint import pprint\nimport timm\nfrom src.utils.metric_learning_loss import *\nfrom src.utils.metrics import *\nfrom src.utils.loss import *\nimport os\nimport torch.nn as nn\nimport pandas as pd\nimport numpy as np\nfrom pdb import set_trace as st\nfrom sklearn.preprocessing import OneHotEncoder\nimport torch.nn.functional as F\nimport pytorch_lightning as pl\nfrom types import MethodType\n\nfrom src.models.resnet3d_csn import *\nfrom src.models.uniformerv2 import *\nfrom src.models.rsna import *\nfrom src.models.layers import AdaptiveConcatPool2d, Flatten\nfrom src.models.ch_mdl_dolg_efficientnet import ChMdlDolgEfficientnet, ArcFaceLossAdaptiveMargin\nfrom src.models.rsna_multi_image import MultiLevelModel2\nfrom src.models.backbones import *\nfrom src.models.group_norm import convert_groupnorm\nfrom src.models.batch_renorm import convert_batchrenorm\nfrom src.models.multi_instance import MultiInstanceModel, MetaMIL, AttentionMILModel, MultiInstanceModelWithWataruAttention\nfrom src.models.resnet import resnet18, resnet34, resnet101, resnet152\nfrom src.models.nextvit import NextVitNet\nfrom src.models.model_4channels import get_attention, get_resnet34, get_attention_inceptionv3\nfrom src.models.vae import VAE, ResNet_VAE\nfrom src.models.model_with_arcface import ArcMarginProduct, AddMarginProduct, ArcMarginProductSubcenter, ArcMarginProductOutCosine, ArcMarginProductSubcenterOutCosine, PudaeArcNet, WithArcface, WhalePrev1stModel, Guie2\nfrom src.models.with_meta_models import WithMetaModel\n\nfrom src.utils.augmentations.strong_aug import *\nfrom src.utils.augmentations.augmentation import *\nfrom src.utils.augmentations.policy_transform import policy_transform\nfrom sklearn.metrics import roc_auc_score, confusion_matrix, mean_squared_error, average_precision_score, accuracy_score\n\n\n################\n\nclass Baseline:\n    def __init__(self):\n        self.gpu = 'small'\n        self.compe = 'rsna'\n        self.batch_size = 16\n        self.grad_accumulations = 1\n        self.lr = 0.0001\n        self.epochs = 20\n        self.resume = False\n        self.seed = 2023\n        self.tta = 1\n        self.model_name = 'convnext_small.fb_in22k_ft_in1k_384'\n        # self.model_name = 'resnet50'\n        self.num_classes = 1\n        self.model = timm.create_model(self.model_name, pretrained=False, num_classes=self.num_classes)\n        self.criterion = torch.nn.BCEWithLogitsLoss()\n        # self.criterion = torch.nn.BCELoss()\n        # self.transform = medical_v1\n        self.transform = kuma_aug\n        self.image_size = 384\n        self.label_features = ['target']\n        self.metric = roc_auc_score # AUC().torch # MultiAP().torch # MultiAUC().torch\n        self.fp16 = True\n        self.optimizer = 'adam'\n        self.scheduler = 'CosineAnnealingWarmRestarts'\n        self.eta_min = 5e-7\n        self.train_by_all_data = False\n        self.early_stop_patience = 1000\n        self.inference = False\n        self.predict_valid = False\n        self.predict_test = False\n        self.logit_to = None\n        self.pretrained_path = None\n        self.sync_batchnorm = True\n        # self.sync_batchnorm = False\n        self.warmup_epochs = -1\n        self.finetune_transform = base_aug_v1\n        self.mixup = False\n        self.arcface = False\n        self.box_crop = None\n        self.predicted_mask_crop = None\n        self.pad_square = False\n        self.resume_epoch = 0\n        self.t_max=30\n        self.save_top_k = 1\n        self.meta_cols = []\n        self.output_features = False\n        self.force_use_model_path_config_when_inf = None\n        self.reset_classifier_when_inf = False\n        self.upsample = None\n        self.in_chans = 3\n        self.add_imsizes_when_inference = [(0, 0)]\n        self.inf_fp16 = False\n        self.distill = False\n        self.reload_dataloaders_every_n_epochs = 0\n        self.tranform_dataset_version = None\n        self.no_trained_model_when_inf = False\n        self.normalize_horiz_orientation = False\n        self.upsample_batch_pos_n = None\n        self.cut_200 = False\n        self.affine_for_gbr = False\n        self.half_dark = False\n        self.crop_by_left_right_line_text = False\n        self.use_wandb = True\n        self.memo = ''\n        self.use_last_ckpt_when_inference = True\n        self.inference_only = False\n        self.valid_df = None\n        self.valid_df_path = None\n\nclass hms_criterion(nn.Module):\n    def __init__(self):\n        super(hms_criterion, self).__init__()\n        # self.criterion = nn.KLDivLoss(reduction=\"batchmean\")\n        self.criterion = nn.KLDivLoss(reduction=\"mean\")\n    def forward(self, logits, targets, weights=None):\n        logits = F.log_softmax(logits, dim=1)\n        return self.criterion(logits, targets)\n\nclass hms_criterion(nn.Module):\n    def __init__(self):\n        super(hms_criterion, self).__init__()\n        self.criterion = nn.KLDivLoss(reduction=\"none\")  # 'none'を使用して個々の損失を保持\n        self.criterion_aux = torch.nn.BCEWithLogitsLoss()\n\n    def forward(self, logits, targets, weights=None):\n        aux = False\n        if logits.shape[1]>6:\n            targets_aux = targets[:, 6:]\n            logits_aux = logits[:, 6:]\n            loss_aux = self.criterion_aux(logits_aux, targets_aux)\n            logits = logits[:, :6]\n            targets = targets[:, :6]\n            aux = True\n\n        logits = F.log_softmax(logits, dim=1)\n        loss = self.criterion(logits, targets)\n        if weights is not None:\n            loss = loss * weights.unsqueeze(1)\n        loss = loss.mean()\n        if aux:\n            loss = loss + loss_aux*0.2\n        return loss\n\ndef kl_divergence(true, preds):\n#     true=true[:,:6]\n#     preds=preds[:,:6]\n    epsilon = 1e-10  # 0による除算を避けるための小さな値\n    preds = F.softmax(preds.float(), dim=1) + epsilon\n    true = true + epsilon\n    kl_div = true * torch.log(true / preds)\n    return - torch.sum(torch.mean(kl_div, dim=0))\n\n\nclass hms_base(Baseline):\n    def __init__(self):\n        super().__init__()\n        self.compe = 'hms'\n        self.predict_valid = False\n        self.predict_test = False\n        self.model_name = 'convnext_small.fb_in22k_ft_in1k_384'\n        self.image_size = (304, 400)\n        # self.model_name = 'swin_base_patch4_window7_224'\n        # self.image_size = 224\n        self.label_features = ['seizure','lpd','gpd','lrda','grda','other']\n#         self.label_features = ['seizure','lpd','gpd','lrda','grda','other','patient_seizure','patient_lpd','patient_gpd','patient_lrda','patient_grda','patient_other']\n        self.num_classes = len(self.label_features)\n        self.in_chans = 3\n        self.model = timm.create_model(self.model_name, pretrained=False, num_classes=self.num_classes, in_chans=self.in_chans)\n        # self.transform = hms_v1\n        self.transform = hms_aug_v12\n        # self.transform = nodoca_aug\n        self.batch_size = 16\n        self.lr = 1e-4\n        self.grad_accumulations = 1\n        self.metric = None\n        self.criterion = hms_criterion()\n        self.use_eeg_spectrograms = False\n        self.use_center_sec = None\n        self.clip_exp = np.exp(8)\n        self.epochs = 8\n        self.mixup = True\n        self.warmup_epochs = -1\n        self.ch3_zero_or_one = False\n        self.spe_and_eeg = False\n        self.use_one_ll = None\n        self.use_last_ckpt_when_inference = True\n        self.mu_std = None\n        self.vote_over9_weight = 1\n        self.metric = kl_divergence\n\nclass hms_swin_base(hms_base):\n    def __init__(self):\n        super().__init__()\n        self.image_size = 384\n        self.use_wandb = False\n        self.use_eeg_spectrograms = True\n        self.mu_std = (9, 3.5)\n        self.transform = hms_aug_v12\n        self.model_name = 'swinv2_large_window12to24_192to384.ms_in22k_ft_in1k'\n        self.model = timm.create_model(self.model_name, pretrained=False, num_classes=self.num_classes, in_chans=self.in_chans)\n        self.batch_size //= 4\n        self.grad_accumulations *= 4\n\nclass hms_chris_fmax30_30sec_8ims_bandpass(hms_swin_base):\n    def __init__(self):\n        super().__init__()\n        self.use_wandb = False\n        self.use_eeg_spectrograms = True\n        self.mu_std = (9, 3.5)\n        self.transform = hms_aug_v12\n        self.model_name = 'swinv2_large_window12to24_192to384.ms_in22k_ft_in1k'\n        self.model = timm.create_model(self.model_name, pretrained=False, num_classes=self.num_classes, in_chans=self.in_chans)\nclass finetune_hms_chris_fmax30_30sec_8ims_bandpass(hms_chris_fmax30_30sec_8ims_bandpass):\n    def __init__(self, fold):\n        super().__init__()\n        self.pretrained_path = f'/kaggle/input/first-fmax30-30sec-8ims-bandpass/last_fold{fold}.ckpt'\n        self.epochs = 7\n        self.lr = 1e-5\n        self.train_df = pd.read_csv('/kaggle/working/train.csv')\n        self.predict_valid = False\n        self.predict_test = False\n\n\nclass hms_chris_fmax30_50sec_16ims_bandpass(hms_swin_base):\n    def __init__(self):\n        super().__init__()\n        self.use_wandb = False\n        self.use_eeg_spectrograms = True\n        self.mu_std = (9, 3.5)\n        self.transform = hms_aug_v12\n        self.model_name = 'swinv2_large_window12to24_192to384.ms_in22k_ft_in1k'\n        self.model = timm.create_model(self.model_name, pretrained=True, num_classes=self.num_classes, in_chans=self.in_chans)\n        self.batch_size //= 2\n        self.grad_accumulations *= 2\n\nclass finetune_hms_chris_fmax30_50sec_16ims_bandpass(hms_chris_fmax30_50sec_16ims_bandpass):\n    def __init__(self, fold):\n        super().__init__()\n        self.pretrained_path = f'/kaggle/input/first-fmax30-50sec-16ims-bandpass/last_fold{fold}.ckpt'\n        self.epochs = 7\n        self.lr = 1e-5\n        self.predict_valid = False\n        self.predict_test = False\n        self.train_df = pd.read_csv('/kaggle/working/train.csv')\n\nclass hms_chris_fmax60_40sec_16ims_bandpass(hms_swin_base):\n    def __init__(self):\n        super().__init__()\n        self.use_wandb = False\n        self.use_eeg_spectrograms = True\n        self.mu_std = (9, 3.5)\n        self.transform = hms_aug_v12\n        self.model_name = 'swinv2_large_window12to24_192to384.ms_in22k_ft_in1k'\n        self.model = timm.create_model(self.model_name, pretrained=True, num_classes=self.num_classes, in_chans=self.in_chans)\n        self.batch_size //= 2\n        self.grad_accumulations *= 2\n\nclass finetune_hms_chris_fmax60_40sec_16ims_bandpass(hms_chris_fmax30_50sec_16ims_bandpass):\n    def __init__(self, fold):\n        super().__init__()\n#         self.pretrained_path = f'/kaggle/input/first-fmax60-30sec-16ims-bandpass/last_fold{fold}.ckpt'\n        self.pretrained_path = f'/kaggle/input/hms-chris-fmax60-30sec-16ims-bandpass/last_fold{fold}.ckpt'\n        \n        self.epochs = 7\n        self.lr = 1e-5\n        self.predict_valid = False\n        self.predict_test = False\n        self.train_df = pd.read_csv('/kaggle/working/train.csv')\n\n\nclass hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg(hms_swin_base):\n    def __init__(self):\n        super().__init__()\n        self.use_wandb = False\n        self.use_eeg_spectrograms = True\n        self.mu_std = (9, 3.5)\n        self.transform = hms_aug_v12\n        self.model_name = 'swinv2_large_window12to24_192to384.ms_in22k_ft_in1k'\n        self.model = timm.create_model(self.model_name, pretrained=True, num_classes=self.num_classes, in_chans=self.in_chans)\n        self.batch_size //= 2\n        self.grad_accumulations *= 2\n        self.spe_and_eeg = True\n        # self.eeg_spe_dir_base = '/kaggle/input/create-fixed-eeg-spec/fmax30_sec30_8ims'\n\nclass finetune_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg(hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg):\n    def __init__(self, fold):\n        super().__init__()\n        self.pretrained_path = f'/kaggle/input/first-fmax30-30sec-8ims-bandpass-spe-and-eeg/last_fold{fold}.ckpt'\n        self.epochs = 7\n        self.lr = 1e-5\n        self.predict_valid = False\n        self.predict_test = False\n        self.train_df = pd.read_csv('/kaggle/working/train.csv')\n\nclass hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg(hms_swin_base):\n    def __init__(self):\n        super().__init__()\n        self.use_wandb = False\n        self.use_eeg_spectrograms = True\n        self.mu_std = (9, 3.5)\n        self.transform = hms_aug_v12\n        self.model_name = 'swinv2_large_window12to24_192to384.ms_in22k_ft_in1k'\n        self.model = timm.create_model(self.model_name, pretrained=True, num_classes=self.num_classes, in_chans=self.in_chans)\n        self.batch_size //= 2\n        self.grad_accumulations *= 2\n        self.spe_and_eeg = True\n        # self.eeg_spe_dir_base = '/kaggle/input/fmax90-sec10-8ims/fmax90_sec10_8ims'\n\nclass finetune_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg(hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg):\n    def __init__(self, fold):\n        super().__init__()\n        self.pretrained_path = f'/kaggle/input/first-fmax90-10sec-8ims-bandpass-spe-and-eeg/last_fold{fold}.ckpt'\n        self.epochs = 7\n        self.lr = 1e-5\n        self.predict_valid = False\n        self.predict_test = False\n        self.train_df = pd.read_csv('/kaggle/working/train.csv')\n\nclass hms_chris_fmax60_50sec_win256_16ims_bandpass_spe_and_eeg(hms_swin_base):\n    def __init__(self):\n        super().__init__()\n        self.use_wandb = False\n        self.use_eeg_spectrograms = True\n        self.mu_std = (9, 3.5)\n        self.transform = hms_aug_v12\n        self.model_name = 'swinv2_large_window12to24_192to384.ms_in22k_ft_in1k'\n        self.model = timm.create_model(self.model_name, pretrained=True, num_classes=self.num_classes, in_chans=self.in_chans)\n        self.batch_size //= 2\n        self.grad_accumulations *= 2\n        self.spe_and_eeg = True\n\nclass finetune_hms_chris_fmax60_50sec_win256_16ims_bandpass_spe_and_eeg(hms_chris_fmax60_50sec_win256_16ims_bandpass_spe_and_eeg):\n    def __init__(self, fold):\n        super().__init__()\n        self.pretrained_path = None\n#         self.lr = 1e-5\n        self.predict_valid = False\n        self.predict_test = False\n        self.train_df = pd.read_csv('/kaggle/working/train.csv')\n        self.epochs = 14\n\n","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:28:44.688348Z","iopub.execute_input":"2024-04-06T07:28:44.688827Z","iopub.status.idle":"2024-04-06T07:28:44.706337Z","shell.execute_reply.started":"2024-04-06T07:28:44.688793Z","shell.execute_reply":"2024-04-06T07:28:44.705438Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%writefile /kaggle/working/hms-pipeline/hms_pipeline/src/lightning/lightning_modules/classification.py\n","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:27:09.870017Z","iopub.execute_input":"2024-04-06T07:27:09.870534Z","iopub.status.idle":"2024-04-06T07:27:09.878499Z","shell.execute_reply.started":"2024-04-06T07:27:09.870505Z","shell.execute_reply":"2024-04-06T07:27:09.877658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile /kaggle/working/hms-pipeline/hms_pipeline/src/lightning/lightning_modules/classification.py\n\nfrom collections import OrderedDict\nimport torch.optim as optim\n\nimport pytorch_lightning as pl\nimport torch\nimport numpy as np\nfrom scipy.special import softmax\nfrom sklearn.metrics import roc_auc_score\nfrom pdb import set_trace as st\nfrom .scheduler_optimizer import get_optimizer, get_scheduler\n\ndef mixup(x: torch.Tensor, y: torch.Tensor, alpha: float = 1.0):\n    assert alpha > 0, \"alpha should be larger than 0\"\n    assert x.size(0) > 1, \"Mixup cannot be applied to a single instance.\"\n\n    lam = np.random.beta(alpha, alpha)\n    rand_index = torch.randperm(x.size()[0])\n    mixed_x = lam * x + (1 - lam) * x[rand_index, :]\n    target_a, target_b = y, y[rand_index]\n    return mixed_x, target_a, target_b, lam\n\ndef mixup_hms(x: torch.Tensor, y: torch.Tensor, w: torch.Tensor, alpha: float = 1.0):\n    assert alpha > 0, \"alpha should be larger than 0\"\n    assert x.size(0) > 1, \"Mixup cannot be applied to a single instance.\"\n\n    lam = np.random.beta(alpha, alpha)\n    rand_index = torch.randperm(x.size()[0])\n    mixed_x = lam * x + (1 - lam) * x[rand_index, :]\n    target_a, target_b = y, y[rand_index]\n    weight_a, weight_b = w, w[rand_index]\n    return mixed_x, target_a, target_b, weight_a, weight_b, lam\n\n\nclass MyLightningModule(pl.LightningModule):\n    def __init__(self, cfg):\n        super(MyLightningModule, self).__init__()\n        self.model = cfg.model\n        # if cfg.pretrained_path is not None:\n        #     self.model.load_state_dict(torch.load(cfg.pretrained_path)['state_dict'])\n        self.cfg = cfg\n        self.validation_step_outputs = []\n\n    def forward(self, x):\n        return self.model(x)\n\n    def training_step(self, batch, batch_nb):\n        images, targets = batch\n\n        if self.cfg.compe == 'hms':\n            images, weights = images\n        if self.cfg.mixup and (torch.rand(1)[0] < 0.5) and (self.cfg.warmup_epochs < self.current_epoch) and (images.size(0) > 1):\n            if self.cfg.compe == 'hms':\n                mix_images, target_a, target_b, weight_a, weight_b, lam = mixup_hms(images, targets, weights, alpha=0.2)\n            else:\n                mix_images, target_a, target_b, lam = mixup(images, targets, alpha=0.5)\n            if self.cfg.arcface:\n                logits = self.model(mix_images, targets)\n            else:\n                logits = self.forward(mix_images)\n                # if self.cfg.distill:\n                #     with torch.no_grad():\n                #         for model_n, (model, weight) in enumerate(zip(self.cfg.teacher_models, [0.2, 0.4, 0.4])):\n                #             if model_n == 0:\n                #                 teacher_preds = model(mix_images)*weight\n                #             else:\n                #                 teacher_preds += model(mix_images)*weight\n            if False:\n                pass\n            # if self.cfg.distill:\n            #     if self.cfg.distill_cancer_only:\n            #         loss = self.cfg.criterion((logits[:, [0]]/self.cfg.distill_temperature).sigmoid(), (teacher_preds[:, [0]]/self.cfg.distill_temperature).sigmoid())\n            #     else:\n            #         loss = self.cfg.criterion((logits/self.cfg.distill_temperature).sigmoid(), (teacher_preds/self.cfg.distill_temperature).sigmoid())\n            #     if self.cfg.use_origin_label:\n            #         if self.cfg.criterion_for_origin_ratio == 0.5:\n            #             loss2 = self.cfg.criterion_for_origin(logits[:, 1:], targets[:, 1:])\n            #         else:\n            #             loss2 = self.cfg.criterion_for_origin(logits, targets)\n            #         loss = loss*self.cfg.criterion_for_origin_ratio + loss2*(1-self.cfg.criterion_for_origin_ratio)\n            else:\n                if self.cfg.compe == 'hms':\n                    try:\n                        loss = self.cfg.criterion(logits, target_a, weight_a) * lam + (1 - lam) * self.cfg.criterion(logits, target_b, weight_b)\n                    except:\n                        loss = self.cfg.criterion(logits, target_a) * lam + (1 - lam) * self.cfg.criterion(logits, target_b)\n                else:\n                    loss = self.cfg.criterion(logits, target_a) * lam + (1 - lam) * self.cfg.criterion(logits, target_b)\n        else:\n            if self.cfg.arcface:\n                logits = self.model(images, targets)\n            else:\n                logits = self.forward(images)\n                # if self.cfg.distill:\n                #     with torch.no_grad():\n                #         for model_n, (model, weight) in enumerate(zip(self.cfg.teacher_models, [0.2, 0.4, 0.4])):\n                #             if model_n == 0:\n                #                 teacher_preds = model(images)*weight\n                #             else:\n                #                 teacher_preds += model(images)*weight\n            if False:\n                pass\n            # if self.cfg.distill:\n            #     if self.cfg.distill_cancer_only:\n            #         loss = self.cfg.criterion((logits[:, [0]]/self.cfg.distill_temperature).sigmoid(), (teacher_preds[:, [0]]/self.cfg.distill_temperature).sigmoid())\n            #     else:\n            #         loss = self.cfg.criterion((logits/self.cfg.distill_temperature).sigmoid(), (teacher_preds/self.cfg.distill_temperature).sigmoid())\n            #     if self.cfg.use_origin_label:\n            #         if self.cfg.criterion_for_origin_ratio == 0.5:\n            #             loss2 = self.cfg.criterion_for_origin(logits[:, 1:], targets[:, 1:])\n            #         else:\n            #             loss2 = self.cfg.criterion_for_origin(logits, targets)\n            #         loss = loss*self.cfg.criterion_for_origin_ratio + loss2*(1-self.cfg.criterion_for_origin_ratio)\n            else:\n                if self.cfg.compe == 'hms':\n                    try:\n                        loss = self.cfg.criterion(logits, targets, weights)\n                    except:\n                        loss = self.cfg.criterion(logits, targets)\n                else:\n                    loss = self.cfg.criterion(logits, targets)\n        self.log(\"train_loss\", loss.item(), on_step=False, on_epoch=True)\n        return loss\n\n    def validation_step(self, batch, batch_nb):\n        images, targets = batch\n        if self.cfg.compe == 'hms':\n            images, weights = images\n        logits = self.forward(images)\n\n        if self.cfg.compe == 'hms':\n            try:\n                loss = self.cfg.criterion(logits, targets, weights)\n            except:\n                loss = self.cfg.criterion(logits, targets)\n        else:\n            loss = self.cfg.criterion(logits, targets)\n        # if isinstance(logits, tuple):\n        #     # feature = logits[1]\n        #     logits = logits[0]\n        preds = logits\n        # preds = logits.sigmoid()\n        output = OrderedDict({\n            \"targets\": targets.detach(), \"preds\": preds.detach(), \"loss\": loss.detach()\n        })\n        self.validation_step_outputs.append(output)\n        return output\n\n\n\n    def on_validation_epoch_end(self):\n        outputs = self.validation_step_outputs\n        d = dict()\n        d[\"epoch\"] = int(self.current_epoch)\n        d[\"v_loss\"] = torch.stack([o[\"loss\"] for o in outputs]).mean().item()\n\n        targets = torch.cat([o[\"targets\"] for o in outputs]).cpu()#.numpy()\n        preds = torch.cat([o[\"preds\"] for o in outputs]).cpu()#.numpy()\n        if self.cfg.metric is None:\n            score = -d['v_loss']\n        elif len(np.unique(targets)) == 1:\n            score = 0\n        else:\n            # score = self.cfg.metric(targets, preds)\n            score = self.cfg.metric(targets, preds)\n\n\n        d[\"val_metric\"] = score\n        print('val metric:', score)\n        # np.save(f'{self.cfg.output_path}/val_preds/fold{self.cfg.fold}/epoch{self.current_epoch}.npy', preds)\n        self.log_dict(d, prog_bar=True, sync_dist=True)\n\n    def configure_optimizers(self):\n        optimizer = get_optimizer(self.cfg)\n        return {\n            \"optimizer\": optimizer,\n            \"lr_scheduler\": {\n                \"scheduler\": get_scheduler(self.cfg, optimizer),\n                \"monitor\": 'val_metric',\n                \"frequency\": 1\n            }\n        }\n\n    # def configure_optimizers(self):\n    #     optimizer = get_optimizer(self.cfg)\n\n    #     # We don't return the lr scheduler because we need to apply it per iteration, not per epoch\n    #     self.lr_scheduler = CosineWarmupScheduler(\n    #         optimizer, warmup=2, max_iters=20\n    #     )\n    #     return optimizer\n\n    # def optimizer_step(self, *args, **kwargs):\n    #     super().optimizer_step(*args, **kwargs)\n    #     self.lr_scheduler.step()  # Step per iteration\n\n\n    # learning rate warm-up\n    # def optimizer_step(self, current_epoch, batch_nb, optimizer, optimizer_idx, closure, on_tpu=False, using_native_amp=False, using_lbfgs=False):\n    #     # warm up lr\n    #     if self.trainer.global_step < 500:\n    #         lr_scale = min(1., float(self.trainer.global_step + 1) / 500.)\n    #         for pg in optimizer.param_groups:\n    #             pg['lr'] = lr_scale * self.hparams.learning_rate\n\n    #     # update params\n    #     optimizer.step(closure=closure)\n\nclass CosineWarmupScheduler(optim.lr_scheduler._LRScheduler):\n    def __init__(self, optimizer, warmup, max_iters):\n        self.warmup = warmup\n        self.max_num_iters = max_iters\n        super().__init__(optimizer)\n\n    def get_lr(self):\n        lr_factor = self.get_lr_factor(epoch=self.last_epoch)\n        return [base_lr * lr_factor for base_lr in self.base_lrs]\n\n    def get_lr_factor(self, epoch):\n        lr_factor = 0.5 * (1 + np.cos(np.pi * epoch / self.max_num_iters))\n        if epoch <= self.warmup:\n            lr_factor *= epoch * 1.0 / self.warmup\n        return lr_factor\n","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:27:09.879778Z","iopub.execute_input":"2024-04-06T07:27:09.880334Z","iopub.status.idle":"2024-04-06T07:27:09.893641Z","shell.execute_reply.started":"2024-04-06T07:27:09.880305Z","shell.execute_reply":"2024-04-06T07:27:09.892801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile /kaggle/working/hms-pipeline/hms_pipeline/src/lightning/data_modules/classification.py\n\nimport numpy as np\nimport pandas as pd\nimport json\nfrom pathlib import Path\nimport pickle\nfrom glob import glob\nimport cv2\nfrom PIL import Image\nimport random\n\nimport torch\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nfrom torchvision import transforms as T\nfrom pdb import set_trace as st\n\nfrom PIL import ImageFile\nImageFile.LOAD_TRUNCATED_IMAGES = True\nimport pytorch_lightning as pl\nfrom .util import *\n# from .mil import MilDataset\nimport os\nimport librosa\nfrom multiprocessing import Pool, cpu_count\n# from pfio.cache import MultiprocessFileCache\n# from monai.transforms import Resize\nfrom albumentations import ReplayCompose\n\ndef sigmoid(x):\n    return 1/(1 + np.exp(-x))\n\nimport pdb\nimport sys\n# ForkedPdb().set_trace()\nclass ForkedPdb(pdb.Pdb):\n    \"\"\"A Pdb subclass that may be used\n    from a forked multiprocessing child\n    \"\"\"\n    def interaction(self, *args, **kwargs):\n        _stdin = sys.stdin\n        try:\n            sys.stdin = open('/dev/stdin')\n            pdb.Pdb.interaction(self, *args, **kwargs)\n        finally:\n            sys.stdin = _stdin\n\ndef pad_to_square(a, wh_ratio=4):\n    if len(a.shape) == 2:\n        a = np.array([a,a,a]).transpose(1,2,0)\n        grayscale = True\n    else:\n        grayscale = False\n\n    \"\"\" Pad an array `a` evenly until it is a square \"\"\"\n    if a.shape[1]>a.shape[0]*wh_ratio: # pad height\n        n_to_add = a.shape[1]/wh_ratio-a.shape[0]\n\n        pad = int(n_to_add//2)\n        # bottom_pad = int(n_to_add-top_pad)\n        a = np.pad(a, [(pad, pad), (0, 0), (0, 0)], mode='constant')\n\n    elif a.shape[0]*wh_ratio>a.shape[1]: # pad width\n        n_to_add = a.shape[0]*wh_ratio-a.shape[1]\n        pad = int(n_to_add//2)\n        # right_pad = int(n_to_add-left_pad)\n        a = np.pad(a, [(0, 0), (pad, pad), (0, 0)], mode='constant')\n    if grayscale:\n        a = a[:,:,0]\n    return a\n\ndef slice_image(img):\n    _, width,ch = img.shape\n    new_width = int(width / 2)\n    left_img = img[:, 0:new_width]\n    right_img = img[:, new_width:]\n    return left_img, right_img\n\ndef mask_pixels(img):\n    img[img > 1] = 1\n    return img\n\ndef count_pixels(img):\n    return np.sum(img)\n\ndef normalize_horiz_orientation(img, reverse=False):\n    left_img, right_img = slice_image(mask_pixels(img.copy()))\n    if reverse:\n        if count_pixels(right_img) <= count_pixels(left_img):\n            return cv2.flip(img, 1)\n    else:\n        if count_pixels(right_img) > count_pixels(left_img):\n            return cv2.flip(img, 1)\n    return img\n\n# resize_transform = A.Compose([A.Resize(height=self.cfg.image_size[0], width=self.cfg.image_size[1], p=1.0)])\ndef load_image(args):\n    path, imsize = args\n    image = cv2.imread(path)[:,:,::-1]\n    # 画像を拡大する場合は、 INTER_LINEARまたはINTER_CUBIC補間を使用することをお勧めします。画像を縮小する場合は、 INTER_AREA補間を使用することをお勧めします。\n    # キュービック補間は計算が複雑であるため、線形補間よりも低速です。ただし、結果の画像の品質は高くなります。\n    return path, cv2.resize(image, (imsize[1], imsize[0]), interpolation=cv2.INTER_AREA)\n\ndef rsna_load_img(path, cfg):\n    strides = cfg.strides\n    image_size = cfg.image_size\n    im = cv2.imread(path)[:,:,0]\n    # im = cv2.resize(im, image_size)\n    img = [im]\n    for stride in strides:\n        prev_path, next_path = rsna_get_prev_next_path_v2(path, stride)\n        # print('\\n',prev_path, path, next_path,'\\n')\n        try:\n            prev = cv2.imread(prev_path)[:,:,0]\n        except:\n            prev = np.zeros(im.shape)\n        # prev = cv2.resize(prev, image_size)\n        img.append(prev)\n        try:\n            next_ = cv2.imread(next_path)[:,:,0]\n        except:\n            next_ = np.zeros(im.shape)\n\n        # next_ = cv2.resize(next_, image_size)\n        img.append(next_)\n    img = np.array(img).transpose(1,2,0)\n    # img = img.astype('float32') # original is uint16\n    return img.astype('uint8')\n\ndef rsna_get_prev_next_path_v2(path, stride):\n    id = path.split('_')[-1]\n    path_base = '_'.join(path.split('_')[:-1])\n    origin_slice = int(id.split('.')[0])\n    prev = origin_slice - stride\n    prev_path = f'{path_base}_{str(prev).zfill(4)}.png'\n    # for i in range(stride):\n    #     if os.path.exists(prev_path):\n    #         break\n    #     else:\n    #         prev+=1\n    #         prev_path = f'{path_base}_{str(prev).zfill(4)}.png'\n\n    next_ = origin_slice + stride\n    next_path = path.replace(path.split('/')[-1], f'{str(next_).zfill(4)}.png')\n    next_path = f'{path_base}_{str(next_).zfill(4)}.png'\n    # for i in range(stride):\n    #     if os.path.exists(next_path):\n    #         break\n    #     else:\n    #         next_-=1\n    #         next_path = f'{path_base}_{str(next_).zfill(4)}.png'\n    return prev_path, next_path\n\n\nclass HMSDataset2(Dataset):\n    def __init__(self, df, transforms, cfg, phase, current_epoch=None):\n        self.transforms = transforms\n        self.paths = df.path.values\n        self.cfg = cfg\n        self.df = df\n        self.phase = phase\n        if phase != 'test':\n            self.labels = df[cfg.label_features].values\n\n        df['w'] = 1\n        self.criterion_weights = df.w.values\n\n    def __len__(self):\n        return self.df.eeg_id.nunique()\n\n    def __getitem__(self, idx):\n        eeg_id = self.df.eeg_id.unique()[idx]\n        idf = self.df[self.df.eeg_id==eeg_id]\n        i = np.array(torch.randperm(len(idf)))[0]\n        row = idf.iloc[i]\n        sec = int(row.eeg_label_offset_seconds)\n        # if eeg_id==3301091280:\n        #     print('='*1000)\n        #     print('3301091280 sec:', sec)\n        #     print('='*1000)\n        path = row.path\n        path = path.replace('.npy', f'_{sec}.npy')\n        image = hms_load_image(self.cfg, path)\n\n        image = self.transforms(image=image)['image']\n        image = torch.FloatTensor(image.transpose((2, 0, 1)))\n        if self.phase == 'test':\n            return image\n\n        image = (image, self.criterion_weights[idx])\n        label = row[self.cfg.label_features]\n        label = label/np.sum(label)\n        # print(label)\n\n        if type(self.cfg.label_features) == list:\n            return image, torch.FloatTensor(label) # multi class\n        else:\n            if str(self.cfg.criterion) == 'CrossEntropyLoss()':\n                return image, label\n            # return image, torch.FloatTensor(label) # multi class\n            return image, label.astype(np.float32)\n\n\nclass HMSDataset(Dataset):\n    def __init__(self, df, transforms, cfg, phase, current_epoch=None):\n        self.transforms = transforms\n        self.paths = df.path.values\n        self.cfg = cfg\n        self.phase = phase\n        self.current_epoch = current_epoch\n        if phase != 'test':\n            self.labels = df[cfg.label_features].values\n        if len(self.cfg.meta_cols) != 0:\n            self.metas = df[cfg.meta_cols].values\n\n    def __len__(self):\n        return len(self.paths)\n\n    def __getitem__(self, idx):\n        path = self.paths[idx]\n        images = hms_preprocess(self.cfg, path)\n        try:\n            ims = []\n            for image in images:\n                image = self.transforms(image=image)['image']\n                image = torch.FloatTensor(image.transpose((2, 0, 1)))\n                ims.append(image)\n            image = torch.stack(ims)\n            # print('image.size():', image.size())\n\n        except Exception as e:\n            print(e, path)\n            raise\n\n        if self.phase == 'test':\n            return image\n\n        label = self.labels[idx]\n        image = (image, label) # dummy weight\n\n        if type(self.cfg.label_features) == list:\n            return image, torch.FloatTensor(label) # multi class\n        else:\n            if str(self.cfg.criterion) == 'CrossEntropyLoss()':\n                return image, label\n            # return image, torch.FloatTensor(label) # multi class\n            return image, label.astype(np.float32)\n\ndef load_eeg_spectrograms_concat(cfg, path):\n    image = np.load(path)\n    if cfg.mu_std is not None:\n        image = (image-cfg.mu_std[0])/cfg.mu_std[1]\n    image = np.concatenate([image[:,:,i] for i in range(image.shape[2])], 0)\n    image = np.array([image, image, image]).transpose((1,2,0))\n    return image\n\n\ndef hms_preprocess_concat(cfg, row, zero_or_one):\n    path = row.path\n    if (cfg.use_eeg_spectrograms) & (not cfg.spe_and_eeg):\n        return load_eeg_spectrograms_concat(cfg, path)\n\n    spe_path = row.spe_path\n    image = pd.read_parquet(spe_path)\n    t = row.spectrogram_label_offset_seconds\n    image = image[(image.time>t) & (image.time<t+600)]\n    image = image.values[:, 1:]\n\n    if cfg.use_center_sec is not None:\n        cut = (300-cfg.use_center_sec)//2\n        image = image[cut:-cut, :]\n        assert image.shape[0] == cfg.use_center_sec\n\n    image = np.clip(image, np.exp(-4), np.exp(8))\n    image = np.log(image)\n\n    ep = 1e-6\n    mu, std = (-0.13223082279419138, 2.368639211554277)\n    image = (image-mu)/(std+ep)\n    image = np.nan_to_num(image, nan=0.0)\n    image = np.array([image, image, image]).transpose((1,2,0))\n\n\n    if cfg.spe_and_eeg:\n        path = row.path\n        add_im = load_eeg_spectrograms_concat(cfg, path)\n        add_im = cv2.resize(add_im, (512, add_im.shape[1])).transpose((1,0,2))\n        # print(add_im.shape)\n        image = cv2.resize(image, (512, 256)).transpose((1,0,2))\n        image = np.concatenate([image, add_im], axis=1)\n    if cfg.use_one_ll is not None:\n        image = image[:, cfg.use_one_ll*100:(cfg.use_one_ll+1)*100]\n    return image\n        \n        \nclass ClassificationDataset(Dataset):\n    def __init__(self, df, transforms, cfg, phase, current_epoch=None):\n        self.transforms = transforms\n        self.paths = df.path.values\n        self.df = df\n        self.cfg = cfg\n        self.phase = phase\n        self.current_epoch = current_epoch\n        if phase != 'test':\n            self.labels = df[cfg.label_features].values\n        if self.cfg.box_crop is not None:\n            self.boxes = df[['x_min', 'y_min', 'x_max', 'y_max']].astype(int).values\n        if len(self.cfg.meta_cols) != 0:\n            self.metas = df[cfg.meta_cols].values\n\n        cache_dir = \"/groups/gca50041/ariyasu/image_cache/\"\n        # self._cache = MultiprocessFileCache(\n        #     len(self), dir=cache_dir, do_pickle=True\n        # )\n        if cfg.affine_for_gbr:\n            self.affine_weights = df.weight.values\n        if self.cfg.crop_by_left_right_line_text:\n            self.x_min_lefts = df.x_min_left.values\n            self.x_max_lefts = df.x_max_left.values\n            self.x_min_rights = df.x_min_right.values\n            self.x_max_rights = df.x_max_right.values\n        if (hasattr(cfg, 'nodoca_select_path_ratio')) and (cfg.nodoca_select_path_ratio):\n            self.paths_list = df.paths.values\n            self.im_qualities_list = df.im_qualities.values\n\n        if cfg.compe in ['hms', 'hms2']:\n            if cfg.ch3_zero_or_one:\n                self.zero_or_one_list = df.zero_or_one.values\n            else:\n                self.zero_or_one_list = range(len(df))\n\n            df['w'] = 1\n            self.criterion_weights = df.w.values\n\n    def __len__(self):\n        return len(self.paths)\n\n    def _read_image(self, path):\n        if not os.path.exists(path):\n            print('not exists:', path)\n            raise\n\n        if getattr(self.cfg, 'strides', False):\n            image = rsna_load_img(path, self.cfg)\n        else:\n            if '.npy' in path:\n                image = np.load(path)\n                if len(image.shape)==2:\n                    image = np.array([image, image, image]).transpose((1,2,0))\n            else:\n                image = cv2.imread(path)[:,:,::-1]\n        return image, 0\n\n    def __getitem__(self, idx):\n        if (self.phase == 'train') and (hasattr(self.cfg, 'nodoca_select_path_ratio')) and (self.cfg.select_path_from_quality_ratio):\n            paths = self.paths_list[idx]\n            im_qualities = self.im_qualities_list[idx]\n            path = random.choices(paths, weights=im_qualities**self.cfg.select_path_from_quality_ratio)\n        else:\n            path = self.paths[idx]\n        reload_dataloader = (self.phase == 'train') and \\\n            (self.cfg.reload_dataloaders_every_n_epochs != 0) and \\\n            (self.current_epoch % self.cfg.reload_dataloaders_every_n_epochs == 0)\n        epoch_path = f'{path}_{self.cfg.tranform_dataset_version}_fold{self.cfg.fold}_epoch{self.current_epoch}.npy'\n        if reload_dataloader & os.path.exists(epoch_path):\n            image = np.load(epoch_path)\n            image = torch.tensor(image)\n        else:\n            image = hms_preprocess_concat(self.cfg, self.df.iloc[idx], self.zero_or_one_list[idx])\n\n            if self.cfg.affine_for_gbr:\n                # raise\n                before = image.mean(0).mean(0)\n                image = image[:,:,::-1]\n                weight = np.array(eval(self.affine_weights[idx]))\n                weight = weight.reshape(3, 4)\n                color = np.concatenate(\n                    [image, np.ones([image.shape[0], image.shape[1], 1])], -1\n                )\n                image = np.einsum(\"ij,hwj->hwi\", weight.astype(float), color.astype(float))\n                # image = np.einsum(\"ij,hwj->hwi\", weight.astype(np.float64), color.astype(np.float64))\n\n                image[image <= 0] = 0\n                image[image >= 255] = 255\n                image = image[:,:,::-1].astype(np.uint8)\n                # after = image.mean(0).mean(0)\n                # print(before, after, '\\n')\n            # else:\n            #     image = image.astype(np.float)\n\n            if self.cfg.box_crop is not None:\n                box = self.boxes[idx]\n                x_min = np.max([box[0]-self.cfg.box_crop, 0])\n                y_min = np.max([box[1]-self.cfg.box_crop, 0])\n                x_max = np.min([box[2]+self.cfg.box_crop, image.shape[1]])\n                y_max = np.min([box[3]+self.cfg.box_crop, image.shape[0]])\n                image = image[int(y_min):int(y_max), int(x_min):int(x_max), :]\n                # cv2.imwrite('tmp.png', image)\n                # print(image.shape)\n\n            if self.cfg.cut_200:\n                image = image[:, 200:-200, :]\n\n            if self.cfg.half_dark:\n                image[image.shape[0]//2:]= 0\n\n            if self.cfg.pad_square:\n                image = pad_to_square(image, self.cfg.image_size[1]/self.cfg.image_size[0])\n\n            # if getattr(self.cfg, 'blacken_all', False):\n            #     image[:, :, :]= 0\n            if getattr(self.cfg, 'blacken_center', False):\n                w = image.shape[1]\n                image[:, w//4:w//4+w//2] = 0\n                # cv2.imwrite('tmp.png', image)\n                # print(image.shape)\n            # if getattr(self.cfg, 'blacken_notcenter', False):\n            #     h = image.shape[1]\n            #     w = image.shape[1]\n            #     image[:, :w//5*2] = 0\n            #     image[:, w//6+w//2:] =0\n            #     image[:h//2] = 0\n            #     image[h//2+h//5*2:] = 0\n            #     # cv2.imwrite('tmp.png', image)\n            #     # print(image.shape)\n\n\n            if self.transforms:\n                try:\n                    if self.cfg.compe == 'hms':\n                        if self.phase == 'train':\n                            if (hasattr(self.cfg, 'concat_im_N')) and (torch.rand(1)[0] < 0.5):\n                                origin_shape = image.shape\n                                one_im_h = image.shape[0]//self.cfg.concat_im_N\n                                image = [image[i*one_im_h:(i+1)*one_im_h, :, :] for i in range(self.cfg.concat_im_N)]\n                                rand_index = torch.randperm(self.cfg.concat_im_N)\n                                image = np.array(image)[rand_index]\n                                image = np.concatenate(image, 0)\n                                assert origin_shape == image.shape\n                                # ForkedPdb().set_trace()\n\n                    image = self.transforms(image=image)['image']\n\n                    if self.cfg.compe == 'hms':\n                        image = torch.FloatTensor(image.transpose((2, 0, 1)))\n                        if self.phase != 'test':\n                            image = (image, self.criterion_weights[idx])\n                    # elif 'Normalize' not in str(self.transforms):\n                    #     mu, std = image.mean(), image.std()\n                    #     image = (image-mu)/std\n                    #     image = image.transpose(2, 0, 1).astype('float16')\n                    #     image = torch.from_numpy(image)\n\n                    # if 'Normalize' not in str(self.transforms):\n                    #     image = (image/255).astype('float16')\n                    #     image = image.transpose(2, 0, 1)\n                    #     image = torch.from_numpy(image)\n                except Exception as e:\n                    print(e)\n                    # print('x_min, x_max, y_min, y_max:', x_min, x_max, y_min, y_max)\n                    # print('error:', path)\n                    # image = rsna_load_img(path, reverse, self.cfg)\n                    # image = self.transforms(image=image)['image']\n                    # ForkedPdb().set_trace()\n                    raise\n\n        if reload_dataloader & (not os.path.exists(epoch_path)):\n            np.save(epoch_path, image.numpy())\n            # print('save!', epoch_path)\n\n        # np.save(f'/home/acc12347av/ml_pipeline/scripts/tmp/{path.split(\"/\")[-1]}.npy', image.numpy())\n        if len(self.cfg.meta_cols) != 0:\n            meta = torch.FloatTensor(self.metas[idx])\n            image = (image, meta)\n\n        if self.phase == 'test':\n            return image\n\n        label = self.labels[idx]\n\n        if type(self.cfg.label_features) == list:\n            return image, torch.FloatTensor(label) # multi class\n        else:\n            if str(self.cfg.criterion) == 'CrossEntropyLoss()':\n                return image, label\n            # return image, torch.FloatTensor(label) # multi class\n            return image, label.astype(np.float32)\n\ndef select_numbers_around_m(N, m, in_chans, stride):\n    results = []\n    for i in range(in_chans):\n        i-=(in_chans-1)//2\n        i*=stride\n        if i < 0:\n            results.append(max(0, m+i))\n        else:\n            results.append(min(N, m+i))\n    # print(m,N,results,iis)\n    # # 範囲を超えないようにmの前後2を取得\n    # if in_chans==5:\n    #     result = [max(0, m-2), max(0, m-1), m, min(N, m+1), min(N, m+2)]\n    # elif in_chans==3:\n    #     result = [max(0, m-2), m, min(N, m+2)]\n    return results\n\ndef get_middle_numbers(N, num_groups=15):\n    # 範囲をグループに分ける\n    bins = np.linspace(0, N, num_groups+1, dtype=int)\n\n    middle_nums = []\n    # 各グループで中央の値を取得\n    for i in range(len(bins)-1):\n        start = bins[i]\n        end = bins[i+1]\n        mid = (start + end) // 2\n        middle_nums.append(mid)\n\n    return middle_nums\n\ndef resize_with_padding(img, target_shape=(256, 256, 3)):\n    # 黒色で初期化された目標shapeの空の画像を作成\n    background = np.zeros(target_shape, dtype=np.uint8)\n\n    # 入力された画像のshapeを取得\n    h, w = img.shape[:2]\n\n    # 入力された画像を背景の左上隅に配置\n    background[:h, :w] = img\n\n    return background\n\nimport random\nimport time\nimport numpy as np\nfrom torch.utils.data import Sampler\nclass InterleavedMaskClassBatchSampler(Sampler):\n    def __init__(self, df, cfg):\n        self.df = df\n        self.batch_size = cfg.batch_size\n        self.indices_by_class = {cls: list(df[df['mask_class'] == cls].index) for cls in df['mask_class'].unique()}\n        for indices in self.indices_by_class.values():\n            np.random.shuffle(indices)\n\n    def __iter__(self):\n        all_classes = list(self.indices_by_class.keys())\n        while len(all_classes) > 0:\n            cls = np.random.choice(all_classes)\n            if len(self.indices_by_class[cls]) >= self.batch_size:\n                print(self.indices_by_class[cls][:self.batch_size])\n                for _ in range(self.batch_size):\n                    yield self.indices_by_class[cls].pop()\n            else:\n                all_classes.remove(cls)\n\n    def __len__(self):\n        return sum(len(indices) for indices in self.indices_by_class.values())\n\nclass InterleavedMaskClassBatchSampler(Sampler):\n    def __init__(self, df, cfg):\n\n        df['tmp_for_batch_sampler'] = list(range(len(df)))\n        self.batch_size = cfg.batch_size\n        self.df = df\n        self.init_indices()\n        self.total_len = len(self.batch_indices_list) * self.batch_size\n        # st()\n        # df[df.tmp_for_batch_sampler.isin(self.batch_indices_list[1])]\n\n    def init_indices(self):\n        chunks = []\n        for c in [0, 1, 34, 4]:\n            cdf = self.df[self.df.mask_class == c]\n            cdf = cdf.sample(len(cdf))\n            lst = cdf.tmp_for_batch_sampler.values.tolist()\n            chunks += [lst[i:i+self.batch_size] for i in range(0, len(lst), self.batch_size) if len(lst[i:i+self.batch_size]) == self.batch_size]\n\n        self.batch_indices_list = random.sample(chunks, len(chunks))\n\n\n    def __iter__(self):\n        for batch_indices in self.batch_indices_list:\n            for idx in batch_indices:\n                yield idx\n\n    def __len__(self):\n        return self.total_len\n\nclass InterleavedMaskClassBatchSamplerBK(Sampler):\n    def __init__(self, df, cfg):\n        self.df = df\n        self.batch_size = cfg.batch_size\n        self.init_indices()\n\n    def init_indices(self):\n        self.indices_by_class = {cls: list(self.df[self.df['mask_class'] == cls].index) for cls in self.df['mask_class'].unique()}\n        for indices in self.indices_by_class.values():\n            np.random.shuffle(indices)\n\n    def __iter__(self):\n        self.init_indices()  # 各エポックの開始時にインデックスを再初期化\n        all_classes = list(self.indices_by_class.keys())\n        while len(all_classes) > 0:\n            cls = np.random.choice(all_classes)\n            while len(self.indices_by_class[cls]) >= self.batch_size:\n                for _ in range(self.batch_size):\n                    yield self.indices_by_class[cls].pop()\n            if len(self.indices_by_class[cls]) < self.batch_size and len(self.indices_by_class[cls]) > 0:\n                for _ in range(len(self.indices_by_class[cls])):\n                    yield self.indices_by_class[cls].pop()\n            all_classes.remove(cls)\n\n    def __len__(self):\n        return sum(len(indices) for indices in self.indices_by_class.values())\n\ndef collate_fn(batch):\n    images, targets= list(zip(*batch))\n    # images = torch.stack(images)\n    # targets = torch.stack(targets)\n    return images, targets\n\nclass ClassificationDatasetMultiImage(Dataset):\n    def __init__(self, df, transforms, cfg, phase, current_epoch=None):\n        self.transforms = transforms\n        self.paths = df.path.values\n        self.cfg = cfg\n        self.phase = phase\n        if phase != 'test':\n            self.labels = df[cfg.label_features].values\n        if hasattr(cfg, 'meta_df'):\n            self.meta_df = cfg.meta_df\n        self.mask_classes = df[['mask_class']].values\n    def __len__(self):\n        return len(self.paths)\n\n    def __getitem__(self, idx):\n        # path = [path for path in self.paths if '10026_42932' in path][0]\n        # print(path)\n        # images = np.load(path, allow_pickle=True)\n        images = np.load(self.paths[idx], allow_pickle=True)\n        if getattr(self.cfg, 'resize_by_class', False):\n            mask_class = self.paths[idx].split('_')[-1].replace('.npy', '')\n            self.cfg.n_slice_per_c = self.cfg.class_n_instance_map[mask_class]\n\n        if hasattr(self.cfg, 'resize_z'):\n            images = images[np.newaxis]\n            resize = Resize((self.cfg.resize_z, images.shape[1], images.shape[2]))\n\n            images = np.array(resize(images)[0])\n            assert images.shape == (self.cfg.resize_z, images.shape[1], images.shape[2])\n\n        if getattr(self.cfg, 'cutx', False):\n            images = images.transpose((2,0,1))\n        elif getattr(self.cfg, 'cuty', False):\n            images = images.transpose((1,0,2))\n\n        if getattr(self.cfg, 'cut_xyz_pad', False):\n            x = images.shape[2]//2\n            y = images.shape[1]//2\n            z = images.shape[0]//2\n            if self.cfg.usex == 0:\n                images = images[:,:,:x+self.cfg.cut_xyz_pad]\n            elif self.cfg.usex == 1:\n                images = images[:,:,max(0, x-self.cfg.cut_xyz_pad):]\n            else: raise\n\n            if self.cfg.usey == 0:\n                images = images[:,:y+self.cfg.cut_xyz_pad,:]\n            elif self.cfg.usey == 1:\n                images = images[:,max(0, y-self.cfg.cut_xyz_pad):,:]\n            else: raise\n\n            if self.cfg.usez == 0:\n                images = images[:z+self.cfg.cut_xyz_pad,:,:]\n            elif self.cfg.usez == 1:\n                images = images[max(0, z-self.cfg.cut_xyz_pad):,:,:]\n            else: raise\n\n        if getattr(self.cfg, 'skip_each_n_slice', False):\n            images = images[[i for i in range(len(images)) if i % self.cfg.skip_each_n_slice == self.cfg.skip_each_n_slice_mod]]\n        total_images_len = len(images)\n\n        if getattr(self.cfg, 'include_black_images', False):\n            if len(images) >= self.cfg.n_slice_per_c:\n                indexes = get_middle_numbers(len(images), num_groups=self.cfg.n_slice_per_c)\n            else:\n                indexes = list(range(len(images)))\n        else:\n            indexes = get_middle_numbers(len(images), num_groups=self.cfg.n_slice_per_c)\n\n        ids_list = []\n        if getattr(self.cfg, 'input_3ch', False):\n            ids_list = indexes\n            all_image_num = self.cfg.n_slice_per_c\n            # print(ids_list)\n        elif getattr(self.cfg, 'equal_sample', False):\n            ids_list = [indexes[i:i+self.cfg.in_chans] for i in range(0, len(indexes), self.cfg.in_chans)]\n            all_image_num = self.cfg.n_slice_per_c//self.cfg.in_chans\n        else:\n            for i in indexes:\n                ids = select_numbers_around_m(len(images)-1, i, self.cfg.in_chans, self.cfg.stride)\n                ids_list.append(ids)\n            all_image_num = self.cfg.n_slice_per_c\n\n        transformed_images = []\n        replay = None\n        for ids in ids_list:\n            image = images[ids]\n\n        # for i in range(self.cfg.n_slice_per_c):\n        #     image = images[i*5:(i+1)*5]\n            if getattr(self.cfg, 'with_mask', False):\n                im_with_mask = []\n                for im in image:\n                    im_with_mask.append(im[0])\n                mask = image[len(ids)//2][1]\n                im_with_mask.append(mask*255)\n                image = np.array(im_with_mask)\n                # ForkedPdb().set_trace()\n                image = image.transpose((1,2,0))\n            elif not getattr(self.cfg, 'input_3ch', False):\n                image = image.transpose((1,2,0))\n            if getattr(self.cfg, 'resize_with_padding', False):\n                long_ = max(image.shape[0], image.shape[1])\n                if long_ > self.cfg.image_size[0]:\n                    image = resize_with_padding(image, (long_, long_, 3))\n                else:\n                    image = resize_with_padding(image, (self.cfg.image_size[0], self.cfg.image_size[0], 3))\n\n            if getattr(self.cfg, 'resize_by_class', False):\n                mask_class = self.paths[idx].split('_')[-1].replace('.npy', '')\n                image_size = self.cfg.class_image_size_map[mask_class]\n                image = cv2.resize(image, (image_size, image_size))\n            if getattr(self.cfg, 'replay_compose', False):\n                if replay is None:\n                    sample = self.transforms(image=image)\n                    replay = sample[\"replay\"]\n                    image = sample['image']\n                else:\n                    image = ReplayCompose.replay(replay, image=image)['image']\n            else:\n                image = self.transforms(image=image)['image']\n            image = image.transpose(2, 0, 1).astype(np.float32) / 255.\n            transformed_images.append(image)\n\n        images = np.stack(transformed_images, 0)\n        images = torch.tensor(images).float()\n        if len(images) < all_image_num:\n            n_instance_current, channel, width, heigth = images.size()\n\n            images = torch.cat(\n                [images, torch.zeros(all_image_num - n_instance_current, channel, width, heigth)]\n            )\n\n        if self.phase == 'train' and random.random() < self.cfg.p_rand_order_v1:\n            indices = torch.randperm(images.size(0))\n            images = images[indices]\n\n        if hasattr(self.cfg, 'meta_df'):\n            image_id = self.paths[idx].split('/')[-1].replace('.npy', '')\n            mdf = self.meta_df[self.meta_df.image_id==image_id]\n            if getattr(self.cfg, 'skip_each_n_slice', False):\n                mdf = mdf.iloc[[i for i in range(len(mdf)) if i % self.cfg.skip_each_n_slice == self.cfg.skip_each_n_slice_mod]]\n\n            assert len(mdf) == total_images_len\n            meta = torch.FloatTensor(mdf[self.cfg.meta_cols].values[indexes])\n            # if image_id == '10005_18667':\n            #     ForkedPdb().set_trace()\n            #     mdf.iloc[indexes][['slice']]\n            images = (images, meta)\n\n        if self.phase == 'test':\n            return images\n\n        label = self.labels[idx]\n        # if self.phase == 'train':\n        #     return images, self.paths[idx]\n\n        if type(self.cfg.label_features) == list:\n            return images, torch.FloatTensor(label) # multi class\n        else:\n            if str(self.cfg.criterion) == 'CrossEntropyLoss()':\n                return images, label\n            # return images, torch.FloatTensor(label) # multi class\n            return images, label.astype(np.float32)\n\n\nimport albumentations as A\nclass ClassificationDataset4classes(Dataset):\n    def __init__(self, df, transforms, cfg, phase, current_epoch=None):\n        self.transforms = transforms\n        self.image_ids = df.image_id.unique()\n        self.df = df\n        self.cfg = cfg\n        self.phase = phase\n        if hasattr(cfg, 'meta_df'):\n            self.meta_df = cfg.meta_df\n\n    def __len__(self):\n        return len(self.image_ids)\n\n    def __getitem__(self, idx):\n        image_id = self.image_ids[idx]\n        # image_id = '10026_42932'\n        idf = self.df[self.df.image_id == image_id]\n\n        # print(image_id)\n        image_4classes = []\n        base_path = idf.path.values[0]\n        rep = base_path.split('/')[-1].split('_')[-1]\n        for class_n, class_ in enumerate([0, 1, 34, 4]):\n            path = base_path.replace(rep, f'{class_}.npy')\n            if os.path.exists(path):\n                images = np.load(path, allow_pickle=True)\n            else:\n                images = np.zeros((15, 256, 256)).astype(np.uint8)\n\n            total_images_len = len(images)\n\n            indexes = get_middle_numbers(len(images), num_groups=self.cfg.n_slice_per_cs[class_n])\n\n            ids_list = []\n            for i in indexes:\n                ids = select_numbers_around_m(len(images)-1, i, self.cfg.in_chans, self.cfg.stride)\n                ids_list.append(ids)\n            all_image_num = self.cfg.n_slice_per_cs[class_n]\n\n            transformed_images = []\n            for ids in ids_list:\n                image = images[ids]\n                image = image.transpose((1,2,0))\n                image = cv2.resize(image, self.cfg.image_sizes[class_n])\n\n                image = self.transforms(image=image)['image']\n                image = image.transpose(2, 0, 1).astype(np.float32) / 255.\n                transformed_images.append(image)\n\n            images = np.stack(transformed_images, 0)\n            images = torch.tensor(images).float()\n            if len(images) < all_image_num:\n                n_instance_current, channel, width, heigth = images.size()\n\n                images = torch.cat(\n                    [images, torch.zeros(all_image_num - n_instance_current, channel, width, heigth)]\n                )\n\n            if self.phase == 'train' and random.random() < self.cfg.p_rand_order_v1:\n                indices = torch.randperm(images.size(0))\n                images = images[indices]\n\n            image_4classes.append(images)\n\n        if self.phase == 'test':\n            return image_4classes\n\n        label = idf[self.cfg.label_features].values[0]\n        if type(self.cfg.label_features) == list:\n            return image_4classes, torch.FloatTensor(label) # multi class\n        else:\n            if str(self.cfg.criterion) == 'CrossEntropyLoss()':\n                return image_4classes, label\n            # return image_4classes, torch.FloatTensor(label) # multi class\n            return image_4classes[0], label.astype(np.float32)\n\nclass ExtravasationDataset(Dataset):\n    def __init__(self, df, transforms, cfg, phase, current_epoch=None):\n        self.transforms = transforms\n        self.paths = df.path.values\n        self.df = df\n        self.patient_ids = df.patient_id.unique()\n        self.cfg = cfg\n        self.phase = phase\n        if phase != 'test':\n            self.labels = df[cfg.label_features].values\n\n    def __len__(self):\n        return len(self.patient_ids)\n\n    def __getitem__(self, idx):\n        patient_id = self.patient_ids[idx]\n        pdf = self.df[self.df.patient_id==patient_id].sort_values('aortic_hu')\n        two_series_images = []\n        for series_n in range(2):\n            images = np.load(pdf.path.values[series_n])\n            if getattr(self.cfg, 'skip_each_n_slice', False):\n                images = images[[i for i in range(len(images)) if i % self.cfg.skip_each_n_slice == self.cfg.skip_each_n_slice_mod]]\n\n            if getattr(self.cfg, 'include_black_images', False):\n                if len(images) >= self.cfg.n_slice_per_c:\n                    indexes = get_middle_numbers(len(images), num_groups=self.cfg.n_slice_per_c)\n                else:\n                    indexes = list(range(len(images)))\n            else:\n                indexes = get_middle_numbers(len(images), num_groups=self.cfg.n_slice_per_c)\n\n            ids_list = []\n            if getattr(self.cfg, 'input_3ch', False):\n                ids_list = indexes\n                all_image_num = self.cfg.n_slice_per_c\n                # print(ids_list)\n            elif getattr(self.cfg, 'equal_sample', False):\n                ids_list = [indexes[i:i+self.cfg.in_chans] for i in range(0, len(indexes), self.cfg.in_chans)]\n                all_image_num = self.cfg.n_slice_per_c//self.cfg.in_chans\n            else:\n                for i in indexes:\n                    ids = select_numbers_around_m(len(images)-1, i, self.cfg.in_chans, self.cfg.stride)\n                    ids_list.append(ids)\n                all_image_num = self.cfg.n_slice_per_c\n\n            transformed_images = []\n            for ids in ids_list:\n                image = images[ids]\n\n            # for i in range(self.cfg.n_slice_per_c):\n            #     image = images[i*5:(i+1)*5]\n                if getattr(self.cfg, 'with_mask', False):\n                    im_with_mask = []\n                    for im in image:\n                        im_with_mask.append(im[0])\n                    mask = image[len(ids)//2][1]\n                    im_with_mask.append(mask*255)\n                    image = np.array(im_with_mask)\n                    # ForkedPdb().set_trace()\n                    image = image.transpose((1,2,0))\n                elif not getattr(self.cfg, 'input_3ch', False):\n                    image = image.transpose((1,2,0))\n                image = self.transforms(image=image)['image']\n                image = image.transpose(2, 0, 1).astype(np.float32) / 255.\n                transformed_images.append(image)\n\n            images = np.stack(transformed_images, 0)\n            images = torch.tensor(images).float()\n            if len(images) < all_image_num:\n                n_instance_current, channel, width, heigth = images.size()\n\n                images = torch.cat(\n                    [images, torch.zeros(all_image_num - n_instance_current, channel, width, heigth)]\n                )\n            if self.phase == 'train' and random.random() < self.cfg.p_rand_order_v1:\n                indices = torch.randperm(images.size(0))\n                images = images[indices]\n\n            two_series_images.append(images)\n        two_series_images = tuple(two_series_images)\n        if self.phase == 'test':\n            return two_series_images\n\n        label = pdf[self.cfg.label_features].values[0]\n\n        if type(self.cfg.label_features) == list:\n            return two_series_images, torch.FloatTensor(label) # multi class\n        else:\n            if str(self.cfg.criterion) == 'CrossEntropyLoss()':\n                return two_series_images, label\n            # return two_series_images, torch.FloatTensor(label) # multi class\n            return two_series_images, label.astype(np.float32)\n\nclass KidneyDataset(Dataset):\n    def __init__(self, df, transforms, cfg, phase, current_epoch=None):\n        self.transforms = transforms\n        self.paths = df.path.values\n        self.df = df\n        self.ids = df.image_id.unique()\n        self.cfg = cfg\n        self.phase = phase\n        if phase != 'test':\n            self.labels = df[cfg.label_features].values\n\n    def __len__(self):\n        return len(self.ids)\n\n    def __getitem__(self, idx):\n        id = self.ids[idx]\n        pdf = self.df[self.df.image_id==id].sort_values('mask_class')\n        two_series_images = []\n        for mask_class in [2,3]:\n            mdf = pdf[pdf.mask_class==mask_class]\n            if len(mdf)==0:\n                images = np.zeros((15, 3, 128, 128))\n                images = torch.tensor(images).float()\n                two_series_images.append(images)\n                continue\n            images = np.load(mdf.path.values[0])\n            if getattr(self.cfg, 'skip_each_n_slice', False):\n                images = images[[i for i in range(len(images)) if i % self.cfg.skip_each_n_slice == self.cfg.skip_each_n_slice_mod]]\n\n            if getattr(self.cfg, 'include_black_images', False):\n                if len(images) >= self.cfg.n_slice_per_c:\n                    indexes = get_middle_numbers(len(images), num_groups=self.cfg.n_slice_per_c)\n                else:\n                    indexes = list(range(len(images)))\n            else:\n                indexes = get_middle_numbers(len(images), num_groups=self.cfg.n_slice_per_c)\n\n            ids_list = []\n            if getattr(self.cfg, 'input_3ch', False):\n                ids_list = indexes\n                all_image_num = self.cfg.n_slice_per_c\n                # print(ids_list)\n            elif getattr(self.cfg, 'equal_sample', False):\n                ids_list = [indexes[i:i+self.cfg.in_chans] for i in range(0, len(indexes), self.cfg.in_chans)]\n                all_image_num = self.cfg.n_slice_per_c//self.cfg.in_chans\n            else:\n                for i in indexes:\n                    ids = select_numbers_around_m(len(images)-1, i, self.cfg.in_chans, self.cfg.stride)\n                    ids_list.append(ids)\n                all_image_num = self.cfg.n_slice_per_c\n\n            transformed_images = []\n            for ids in ids_list:\n                image = images[ids]\n\n            # for i in range(self.cfg.n_slice_per_c):\n            #     image = images[i*5:(i+1)*5]\n                if getattr(self.cfg, 'with_mask', False):\n                    im_with_mask = []\n                    for im in image:\n                        im_with_mask.append(im[0])\n                    mask = image[len(ids)//2][1]\n                    im_with_mask.append(mask*255)\n                    image = np.array(im_with_mask)\n                    # ForkedPdb().set_trace()\n                    image = image.transpose((1,2,0))\n                elif not getattr(self.cfg, 'input_3ch', False):\n                    image = image.transpose((1,2,0))\n                image = self.transforms(image=image)['image']\n                image = image.transpose(2, 0, 1).astype(np.float32) / 255.\n                transformed_images.append(image)\n\n            images = np.stack(transformed_images, 0)\n            images = torch.tensor(images).float()\n            if len(images) < all_image_num:\n                n_instance_current, channel, width, heigth = images.size()\n\n                images = torch.cat(\n                    [images, torch.zeros(all_image_num - n_instance_current, channel, width, heigth)]\n                )\n            if self.phase == 'train' and random.random() < self.cfg.p_rand_order_v1:\n                indices = torch.randperm(images.size(0))\n                images = images[indices]\n\n            two_series_images.append(images)\n        two_series_images = tuple(two_series_images)\n        if self.phase == 'test':\n            return two_series_images\n\n        label = pdf[self.cfg.label_features].values[0]\n\n        if type(self.cfg.label_features) == list:\n            return two_series_images, torch.FloatTensor(label) # multi class\n        else:\n            if str(self.cfg.criterion) == 'CrossEntropyLoss()':\n                return two_series_images, label\n            # return two_series_images, torch.FloatTensor(label) # multi class\n            return two_series_images, label.astype(np.float32)\n\n\nclass ClassificationDatasetMultiImage2nd(Dataset):\n    def __init__(self, df, transforms, cfg, phase, current_epoch=None):\n\n        self.transforms = transforms\n        self.paths = df.path.values\n        self.cfg = cfg\n        self.phase = phase\n        if phase != 'test':\n            self.labels = df[cfg.label_features].values\n\n    def __len__(self):\n        return len(self.paths)\n\n    def __getitem__(self, idx):\n        image = np.load(self.paths[idx])\n        if len(image) >= self.cfg.n_slice_per_c:\n            indexes = get_middle_numbers(len(image), num_groups=self.cfg.n_slice_per_c)\n            image = image[indexes]\n        else:\n            image = np.concatenate([image, np.zeros((self.cfg.n_slice_per_c-len(image), image.shape[1], image.shape[2]))])\n\n        image = image.transpose((1,2,0)).astype(np.float32)\n        image = self.transforms(image=image)['image']\n        image = image.transpose(2, 0, 1).astype(np.float32) / 255.\n\n        image = torch.tensor(image).float()\n        if self.phase == 'test':\n            return image\n\n        if self.phase == 'train' and random.random() < self.cfg.p_rand_order_v1:\n            indices = torch.randperm(image.size(0))\n            image = image[indices]\n\n        label = self.labels[idx]\n\n        if type(self.cfg.label_features) == list:\n            return image, torch.FloatTensor(label) # multi class\n        else:\n            if str(self.cfg.criterion) == 'CrossEntropyLoss()':\n                return image, label\n            # return image, torch.FloatTensor(label) # multi class\n            return image, label.astype(np.float32)\n\n\ndef worker_init_fn(worker_id):\n    np.random.seed(np.random.get_state()[1][0] + worker_id)\n\n\ndef get_dataset_class(cfg):\n    if cfg.compe == 'nodoca':\n        claz = MilDataset\n    elif (hasattr(cfg, 'hms_random_sample')) and (cfg.hms_random_sample):\n        claz = HMSDataset2\n    elif (hasattr(cfg, 'hms_multi_image_model')) and (cfg.hms_multi_image_model):\n        claz = HMSDataset\n    elif getattr(cfg, 'multi_image_4classes', False):\n        claz = ClassificationDataset4classes\n    elif getattr(cfg, 'multi_image_extravasation', False):\n        claz = ExtravasationDataset\n    elif getattr(cfg, 'multi_image_kidney', False):\n        claz = KidneyDataset\n    elif getattr(cfg, 'multi_image', False):\n        claz = ClassificationDatasetMultiImage\n    else:\n        claz = ClassificationDataset\n    return claz\n\nclass MyDataModule(pl.LightningDataModule):\n    def __init__(self, cfg):\n        super().__init__()\n        self.cfg = cfg\n\n    # 必ず呼び出される関数\n    def setup(self, stage):\n        pass\n\n\n    # Trainer.fit() 時に呼び出される\n    def train_dataloader(self):\n        if self.cfg.train_by_all_data:\n            tr = self.cfg.train_df\n        else:\n            tr = self.cfg.train_df[self.cfg.train_df.fold != self.cfg.fold]\n        self.tr = tr\n        if self.cfg.upsample is not None:\n            assert type(self.cfg.upsample) == int\n            origin_len = len(tr)\n            dfs = [tr]\n            for col in self.cfg.label_features:\n                for _ in range(self.cfg.upsample):\n                    dfs.append(tr[tr[col]==1])\n            tr = pd.concat(dfs)\n            print(f'upsample, len: {origin_len} -> {len(tr)}')\n\n        print('len(train):', len(tr))\n        claz = get_dataset_class(self.cfg)\n        if getattr(self.cfg, 'use_custom_sampler', False):\n            tr = tr.reset_index(drop=True)\n        train_ds = claz(\n            df=tr,\n            transforms=self.cfg.transform['train'],\n            cfg=self.cfg,\n            phase='train',\n            current_epoch=self.trainer.current_epoch,\n        )\n        if getattr(self.cfg, 'use_custom_sampler', False):\n            return DataLoader(train_ds, batch_size=self.cfg.batch_size,\n                num_workers=self.cfg.n_cpu, worker_init_fn=worker_init_fn,\n                sampler=InterleavedMaskClassBatchSampler(tr, self.cfg))\n\n        else:\n            return DataLoader(train_ds, batch_size=self.cfg.batch_size, pin_memory=True, shuffle=True, drop_last=True,\n                num_workers=self.cfg.n_cpu, worker_init_fn=worker_init_fn)\n\n    # Trainer.fit() 時に呼び出される\n    def val_dataloader(self):\n        val = get_val(self.cfg)\n        self.val = val\n\n        print('len(valid):', len(val))\n        claz = get_dataset_class(self.cfg)\n\n        valid_ds = claz(\n            df=val,\n            transforms=self.cfg.transform['val'],\n            cfg=self.cfg,\n            phase='valid'\n        )\n\n        return DataLoader(valid_ds, batch_size=self.cfg.batch_size, pin_memory=True, shuffle=False, drop_last=False,\n                          num_workers=self.cfg.n_cpu, worker_init_fn=worker_init_fn)\n","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:27:09.895302Z","iopub.execute_input":"2024-04-06T07:27:09.89556Z","iopub.status.idle":"2024-04-06T07:27:09.934048Z","shell.execute_reply.started":"2024-04-06T07:27:09.895538Z","shell.execute_reply":"2024-04-06T07:27:09.933211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile /kaggle/working/hms-pipeline/hms_pipeline/src/models/resnet3d_csn.py\n\n# test","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:27:09.93504Z","iopub.execute_input":"2024-04-06T07:27:09.93529Z","iopub.status.idle":"2024-04-06T07:27:09.946878Z","shell.execute_reply.started":"2024-04-06T07:27:09.935269Z","shell.execute_reply":"2024-04-06T07:27:09.945937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile /kaggle/working/hms-pipeline/hms_pipeline/src/models/uniformerv2.py\n\n# test","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:27:09.948064Z","iopub.execute_input":"2024-04-06T07:27:09.94885Z","iopub.status.idle":"2024-04-06T07:27:09.956165Z","shell.execute_reply.started":"2024-04-06T07:27:09.948826Z","shell.execute_reply":"2024-04-06T07:27:09.955374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile /kaggle/working/hms-pipeline/hms_pipeline/src/models/rsna.py\n\n# test","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:27:09.957331Z","iopub.execute_input":"2024-04-06T07:27:09.957878Z","iopub.status.idle":"2024-04-06T07:27:09.965276Z","shell.execute_reply.started":"2024-04-06T07:27:09.957848Z","shell.execute_reply":"2024-04-06T07:27:09.964205Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile /kaggle/working/hms-pipeline/hms_pipeline/src/models/senet.py\n\n\"\"\"\nhttps://github.com/pytorch/vision/blob/master/torchvision/models/resnet.py\n\"\"\"\nfrom __future__ import print_function, division, absolute_import\nfrom collections import OrderedDict\nimport math\n\nimport torch.nn as nn\nfrom torch.utils import model_zoo\n\n__all__ = ['SENet', 'senet154', 'se_resnet50', 'se_resnet101', 'se_resnet152',\n           'se_resnext50_32x4d', 'se_resnext101_32x4d']\n\npretrained_settings = {\n    'senet154': {\n        'imagenet': {\n            'url': 'http://data.lip6.fr/cadene/pretrainedmodels/senet154-c7b49a05.pth',\n            'input_space': 'RGB',\n            'input_size': [3, 224, 224],\n            'input_range': [0, 1],\n            'mean': [0.485, 0.456, 0.406],\n            'std': [0.229, 0.224, 0.225],\n            'num_classes': 1000\n        }\n    },\n    'se_resnet50': {\n        'imagenet': {\n            'url': 'http://data.lip6.fr/cadene/pretrainedmodels/se_resnet50-ce0d4300.pth',\n            'input_space': 'RGB',\n            'input_size': [3, 224, 224],\n            'input_range': [0, 1],\n            'mean': [0.485, 0.456, 0.406],\n            'std': [0.229, 0.224, 0.225],\n            'num_classes': 1000\n        }\n    },\n    'se_resnet101': {\n        'imagenet': {\n            'url': 'http://data.lip6.fr/cadene/pretrainedmodels/se_resnet101-7e38fcc6.pth',\n            'input_space': 'RGB',\n            'input_size': [3, 224, 224],\n            'input_range': [0, 1],\n            'mean': [0.485, 0.456, 0.406],\n            'std': [0.229, 0.224, 0.225],\n            'num_classes': 1000\n        }\n    },\n    'se_resnet152': {\n        'imagenet': {\n            'url': 'http://data.lip6.fr/cadene/pretrainedmodels/se_resnet152-d17c99b7.pth',\n            'input_space': 'RGB',\n            'input_size': [3, 224, 224],\n            'input_range': [0, 1],\n            'mean': [0.485, 0.456, 0.406],\n            'std': [0.229, 0.224, 0.225],\n            'num_classes': 1000\n        }\n    },\n    'se_resnext50_32x4d': {\n        'imagenet': {\n            'url': 'http://data.lip6.fr/cadene/pretrainedmodels/se_resnext50_32x4d-a260b3a4.pth',\n            'input_space': 'RGB',\n            'input_size': [3, 224, 224],\n            'input_range': [0, 1],\n            'mean': [0.485, 0.456, 0.406],\n            'std': [0.229, 0.224, 0.225],\n            'num_classes': 1000\n        }\n    },\n    'se_resnext101_32x4d': {\n        'imagenet': {\n            'url': 'http://data.lip6.fr/cadene/pretrainedmodels/se_resnext101_32x4d-3b2fe3d8.pth',\n            'input_space': 'RGB',\n            'input_size': [3, 224, 224],\n            'input_range': [0, 1],\n            'mean': [0.485, 0.456, 0.406],\n            'std': [0.229, 0.224, 0.225],\n            'num_classes': 1000\n        }\n    },\n}\n\n\nclass SEModule(nn.Module):\n\n    def __init__(self, channels, reduction):\n        super(SEModule, self).__init__()\n        self.avg_pool = nn.AdaptiveAvgPool2d(1)\n        self.fc1 = nn.Conv2d(channels, channels // reduction, kernel_size=1,\n                             padding=0)\n        self.relu = nn.ReLU(inplace=True)\n        self.fc2 = nn.Conv2d(channels // reduction, channels, kernel_size=1,\n                             padding=0)\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        module_input = x\n        x = self.avg_pool(x)\n        x = self.fc1(x)\n        x = self.relu(x)\n        x = self.fc2(x)\n        x = self.sigmoid(x)\n        return module_input * x\n\n\nclass Bottleneck(nn.Module):\n    \"\"\"\n    Base class for bottlenecks that implements `forward()` method.\n    \"\"\"\n\n    def forward(self, x):\n        residual = x\n\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n\n        out = self.conv2(out)\n        out = self.bn2(out)\n        out = self.relu(out)\n\n        out = self.conv3(out)\n        out = self.bn3(out)\n\n        if self.downsample is not None:\n            residual = self.downsample(x)\n\n        out = self.se_module(out) + residual\n        out = self.relu(out)\n\n        return out\n\n\nclass SEBottleneck(Bottleneck):\n    \"\"\"\n    Bottleneck for SENet154.\n    \"\"\"\n    expansion = 4\n\n    def __init__(self, inplanes, planes, groups, reduction, stride=1,\n                 downsample=None):\n        super(SEBottleneck, self).__init__()\n        self.conv1 = nn.Conv2d(inplanes, planes * 2, kernel_size=1, bias=False)\n        self.bn1 = nn.BatchNorm2d(planes * 2)\n        self.conv2 = nn.Conv2d(planes * 2, planes * 4, kernel_size=3,\n                               stride=stride, padding=1, groups=groups,\n                               bias=False)\n        self.bn2 = nn.BatchNorm2d(planes * 4)\n        self.conv3 = nn.Conv2d(planes * 4, planes * 4, kernel_size=1,\n                               bias=False)\n        self.bn3 = nn.BatchNorm2d(planes * 4)\n        self.relu = nn.ReLU(inplace=True)\n        self.se_module = SEModule(planes * 4, reduction=reduction)\n        self.downsample = downsample\n        self.stride = stride\n\n\nclass SEResNetBottleneck(Bottleneck):\n    \"\"\"\n    ResNet bottleneck with a Squeeze-and-Excitation module. It follows Caffe\n    implementation and uses `stride=stride` in `conv1` and not in `conv2`\n    (the latter is used in the torchvision implementation of ResNet).\n    \"\"\"\n    expansion = 4\n\n    def __init__(self, inplanes, planes, groups, reduction, stride=1,\n                 downsample=None):\n        super(SEResNetBottleneck, self).__init__()\n        self.conv1 = nn.Conv2d(inplanes, planes, kernel_size=1, bias=False,\n                               stride=stride)\n        self.bn1 = nn.BatchNorm2d(planes)\n        self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, padding=1,\n                               groups=groups, bias=False)\n        self.bn2 = nn.BatchNorm2d(planes)\n        self.conv3 = nn.Conv2d(planes, planes * 4, kernel_size=1, bias=False)\n        self.bn3 = nn.BatchNorm2d(planes * 4)\n        self.relu = nn.ReLU(inplace=True)\n        self.se_module = SEModule(planes * 4, reduction=reduction)\n        self.downsample = downsample\n        self.stride = stride\n\n\nclass SEResNeXtBottleneck(Bottleneck):\n    \"\"\"\n    ResNeXt bottleneck type C with a Squeeze-and-Excitation module.\n    \"\"\"\n    expansion = 4\n\n    def __init__(self, inplanes, planes, groups, reduction, stride=1,\n                 downsample=None, base_width=4):\n        super(SEResNeXtBottleneck, self).__init__()\n        width = math.floor(planes * (base_width / 64)) * groups\n        self.conv1 = nn.Conv2d(inplanes, width, kernel_size=1, bias=False,\n                               stride=1)\n        self.bn1 = nn.BatchNorm2d(width)\n        self.conv2 = nn.Conv2d(width, width, kernel_size=3, stride=stride,\n                               padding=1, groups=groups, bias=False)\n        self.bn2 = nn.BatchNorm2d(width)\n        self.conv3 = nn.Conv2d(width, planes * 4, kernel_size=1, bias=False)\n        self.bn3 = nn.BatchNorm2d(planes * 4)\n        self.relu = nn.ReLU(inplace=True)\n        self.se_module = SEModule(planes * 4, reduction=reduction)\n        self.downsample = downsample\n        self.stride = stride\n\n\nclass SENet(nn.Module):\n\n    def __init__(self, block, layers, groups, reduction, dropout_p=0.2,\n                 inplanes=128, input_3x3=True, downsample_kernel_size=3,\n                 downsample_padding=1, num_classes=1000):\n        \"\"\"\n        Parameters\n        ----------\n        block (nn.Module): Bottleneck class.\n            - For SENet154: SEBottleneck\n            - For SE-ResNet models: SEResNetBottleneck\n            - For SE-ResNeXt models:  SEResNeXtBottleneck\n        layers (list of ints): Number of residual blocks for 4 layers of the\n            network (layer1...layer4).\n        groups (int): Number of groups for the 3x3 convolution in each\n            bottleneck block.\n            - For SENet154: 64\n            - For SE-ResNet models: 1\n            - For SE-ResNeXt models:  32\n        reduction (int): Reduction ratio for Squeeze-and-Excitation modules.\n            - For all models: 16\n        dropout_p (float or None): Drop probability for the Dropout layer.\n            If `None` the Dropout layer is not used.\n            - For SENet154: 0.2\n            - For SE-ResNet models: None\n            - For SE-ResNeXt models: None\n        inplanes (int):  Number of input channels for layer1.\n            - For SENet154: 128\n            - For SE-ResNet models: 64\n            - For SE-ResNeXt models: 64\n        input_3x3 (bool): If `True`, use three 3x3 convolutions instead of\n            a single 7x7 convolution in layer0.\n            - For SENet154: True\n            - For SE-ResNet models: False\n            - For SE-ResNeXt models: False\n        downsample_kernel_size (int): Kernel size for downsampling convolutions\n            in layer2, layer3 and layer4.\n            - For SENet154: 3\n            - For SE-ResNet models: 1\n            - For SE-ResNeXt models: 1\n        downsample_padding (int): Padding for downsampling convolutions in\n            layer2, layer3 and layer4.\n            - For SENet154: 1\n            - For SE-ResNet models: 0\n            - For SE-ResNeXt models: 0\n        num_classes (int): Number of outputs in `last_linear` layer.\n            - For all models: 1000\n        \"\"\"\n        super(SENet, self).__init__()\n        self.inplanes = inplanes\n        if input_3x3:\n            layer0_modules = [\n                ('conv1', nn.Conv2d(3, 64, 3, stride=2, padding=1,\n                                    bias=False)),\n                ('bn1', nn.BatchNorm2d(64)),\n                ('relu1', nn.ReLU(inplace=True)),\n                ('conv2', nn.Conv2d(64, 64, 3, stride=1, padding=1,\n                                    bias=False)),\n                ('bn2', nn.BatchNorm2d(64)),\n                ('relu2', nn.ReLU(inplace=True)),\n                ('conv3', nn.Conv2d(64, inplanes, 3, stride=1, padding=1,\n                                    bias=False)),\n                ('bn3', nn.BatchNorm2d(inplanes)),\n                ('relu3', nn.ReLU(inplace=True)),\n            ]\n        else:\n            layer0_modules = [\n                ('conv1', nn.Conv2d(3, inplanes, kernel_size=7, stride=2,\n                                    padding=3, bias=False)),\n                ('bn1', nn.BatchNorm2d(inplanes)),\n                ('relu1', nn.ReLU(inplace=True)),\n            ]\n        # To preserve compatibility with Caffe weights `ceil_mode=True`\n        # is used instead of `padding=1`.\n        layer0_modules.append(('pool', nn.MaxPool2d(3, stride=2,\n                                                    ceil_mode=True)))\n        self.layer0 = nn.Sequential(OrderedDict(layer0_modules))\n        self.layer1 = self._make_layer(\n            block,\n            planes=64,\n            blocks=layers[0],\n            groups=groups,\n            reduction=reduction,\n            downsample_kernel_size=1,\n            downsample_padding=0\n        )\n        self.layer2 = self._make_layer(\n            block,\n            planes=128,\n            blocks=layers[1],\n            stride=2,\n            groups=groups,\n            reduction=reduction,\n            downsample_kernel_size=downsample_kernel_size,\n            downsample_padding=downsample_padding\n        )\n        self.layer3 = self._make_layer(\n            block,\n            planes=256,\n            blocks=layers[2],\n            stride=2,\n            groups=groups,\n            reduction=reduction,\n            downsample_kernel_size=downsample_kernel_size,\n            downsample_padding=downsample_padding\n        )\n        self.layer4 = self._make_layer(\n            block,\n            planes=512,\n            blocks=layers[3],\n            stride=2,\n            groups=groups,\n            reduction=reduction,\n            downsample_kernel_size=downsample_kernel_size,\n            downsample_padding=downsample_padding\n        )\n        # self.avg_pool = nn.AvgPool2d(7, stride=1)\n        self.avg_pool = nn.AdaptiveAvgPool2d((1, 1))\n        self.dropout = nn.Dropout(dropout_p) if dropout_p is not None else None\n        self.last_linear = nn.Linear(512 * block.expansion, num_classes)\n\n    def _make_layer(self, block, planes, blocks, groups, reduction, stride=1,\n                    downsample_kernel_size=1, downsample_padding=0):\n        downsample = None\n        if stride != 1 or self.inplanes != planes * block.expansion:\n            downsample = nn.Sequential(\n                nn.Conv2d(self.inplanes, planes * block.expansion,\n                          kernel_size=downsample_kernel_size, stride=stride,\n                          padding=downsample_padding, bias=False),\n                nn.BatchNorm2d(planes * block.expansion),\n            )\n\n        layers = []\n        layers.append(block(self.inplanes, planes, groups, reduction, stride,\n                            downsample))\n        self.inplanes = planes * block.expansion\n        for i in range(1, blocks):\n            layers.append(block(self.inplanes, planes, groups, reduction))\n\n        return nn.Sequential(*layers)\n\n    def features(self, x):\n        x = self.layer0(x)\n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n        return x\n\n    def logits(self, x):\n        x = self.avg_pool(x)\n        if self.dropout is not None:\n            x = self.dropout(x)\n        x = x.view(x.size(0), -1)\n        x = self.last_linear(x)\n        return x\n\n    def forward(self, x):\n        x = self.features(x)\n        x = self.logits(x)\n        return x\n\n\ndef initialize_pretrained_model(model, num_classes, settings):\n    assert num_classes == settings['num_classes'], \\\n        'num_classes should be {}, but is {}'.format(\n            settings['num_classes'], num_classes)\n    # model.load_state_dict(model_zoo.load_url(settings['url']))\n    model.input_space = settings['input_space']\n    model.input_size = settings['input_size']\n    model.input_range = settings['input_range']\n    model.mean = settings['mean']\n    model.std = settings['std']\n\n\ndef senet154(num_classes=1000, pretrained='imagenet'):\n    model = SENet(SEBottleneck, [3, 8, 36, 3], groups=64, reduction=16,\n                  dropout_p=0.2, num_classes=num_classes)\n    if pretrained is not None:\n        settings = pretrained_settings['senet154'][pretrained]\n        initialize_pretrained_model(model, num_classes, settings)\n    return model\n\n\ndef se_resnet50(num_classes=1000, pretrained='imagenet'):\n    model = SENet(SEResNetBottleneck, [3, 4, 6, 3], groups=1, reduction=16,\n                  dropout_p=None, inplanes=64, input_3x3=False,\n                  downsample_kernel_size=1, downsample_padding=0,\n                  num_classes=num_classes)\n    if pretrained is not None:\n        settings = pretrained_settings['se_resnet50'][pretrained]\n        initialize_pretrained_model(model, num_classes, settings)\n    return model\n\n\ndef se_resnet101(num_classes=1000, pretrained='imagenet'):\n    model = SENet(SEResNetBottleneck, [3, 4, 23, 3], groups=1, reduction=16,\n                  dropout_p=None, inplanes=64, input_3x3=False,\n                  downsample_kernel_size=1, downsample_padding=0,\n                  num_classes=num_classes)\n    if pretrained is not None:\n        settings = pretrained_settings['se_resnet101'][pretrained]\n        initialize_pretrained_model(model, num_classes, settings)\n    return model\n\n\ndef se_resnet152(num_classes=1000, pretrained='imagenet'):\n    model = SENet(SEResNetBottleneck, [3, 8, 36, 3], groups=1, reduction=16,\n                  dropout_p=None, inplanes=64, input_3x3=False,\n                  downsample_kernel_size=1, downsample_padding=0,\n                  num_classes=num_classes)\n    if pretrained is not None:\n        settings = pretrained_settings['se_resnet152'][pretrained]\n        initialize_pretrained_model(model, num_classes, settings)\n    return model\n\n\ndef se_resnext50_32x4d(num_classes=1000, pretrained='imagenet'):\n    model = SENet(SEResNeXtBottleneck, [3, 4, 6, 3], groups=32, reduction=16,\n                  dropout_p=None, inplanes=64, input_3x3=False,\n                  downsample_kernel_size=1, downsample_padding=0,\n                  num_classes=num_classes)\n    if pretrained is not None:\n        settings = pretrained_settings['se_resnext50_32x4d'][pretrained]\n        initialize_pretrained_model(model, num_classes, settings)\n    return model\n\n\ndef se_resnext101_32x4d(num_classes=1000, pretrained='imagenet'):\n    model = SENet(SEResNeXtBottleneck, [3, 4, 23, 3], groups=32, reduction=16,\n                  dropout_p=None, inplanes=64, input_3x3=False,\n                  downsample_kernel_size=1, downsample_padding=0,\n                  num_classes=num_classes)\n    if pretrained is not None:\n        settings = pretrained_settings['se_resnext101_32x4d'][pretrained]\n        initialize_pretrained_model(model, num_classes, settings)\n    return model\n","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:27:09.966632Z","iopub.execute_input":"2024-04-06T07:27:09.966889Z","iopub.status.idle":"2024-04-06T07:27:09.983537Z","shell.execute_reply.started":"2024-04-06T07:27:09.966869Z","shell.execute_reply":"2024-04-06T07:27:09.982656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%cd /kaggle/working/hms-pipeline/hms_pipeline/scripts","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:29:02.092661Z","iopub.execute_input":"2024-04-06T07:29:02.093102Z","iopub.status.idle":"2024-04-06T07:29:02.09979Z","shell.execute_reply.started":"2024-04-06T07:29:02.093069Z","shell.execute_reply":"2024-04-06T07:29:02.098755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\ntrain_test = 'train'\nfold_df = pd.read_csv('/kaggle/input/hms-data/train_unique_egg_v2.csv')\ndf = pd.read_csv('/kaggle/input/hms-data/train_with_group_in_eeg.csv')\n\ndf['n'] = df[['seizure_vote','lpd_vote','gpd_vote','lrda_vote','grda_vote','other_vote']].sum(1)\n\ndf = df[df.n>8]\ndf = df.merge(fold_df[['patient_id', 'fold']].drop_duplicates('patient_id'), on='patient_id')\ndf['spe_path'] = f'/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms/'+df.spectrogram_id.astype(str)+'.parquet'\n\ndfs = []\nfor i, idf in df.groupby(['eeg_id', 'group_in_eeg']):\n    sec = (idf.spectrogram_label_offset_seconds.min() + idf.spectrogram_label_offset_seconds.max())//2\n    idf['spectrogram_label_offset_seconds'] = sec\n    dfs.append(idf.iloc[:1])\ndf = pd.concat(dfs)  \ndf[['seizure','lpd','gpd','lrda','grda','other']] = df[['seizure_vote','lpd_vote','gpd_vote','lrda_vote','grda_vote','other_vote']] / np.array([df[['seizure_vote','lpd_vote','gpd_vote','lrda_vote','grda_vote','other_vote']].values.sum(1).tolist()]*6).T","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:29:03.115668Z","iopub.execute_input":"2024-04-06T07:29:03.11611Z","iopub.status.idle":"2024-04-06T07:29:09.396179Z","shell.execute_reply.started":"2024-04-06T07:29:03.11608Z","shell.execute_reply":"2024-04-06T07:29:09.395405Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dfs = []\nfor id, idf in df.groupby('patient_id'):\n    idf[['patient_seizure','patient_lpd','patient_gpd','patient_lrda','patient_grda','patient_other']] = idf[['seizure','lpd','gpd','lrda','grda','other']].mean(0)\n    dfs.append(idf.iloc[:1])\npdf = pd.concat(dfs)    \n","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:29:09.397774Z","iopub.execute_input":"2024-04-06T07:29:09.398056Z","iopub.status.idle":"2024-04-06T07:29:13.085868Z","shell.execute_reply.started":"2024-04-06T07:29:09.398032Z","shell.execute_reply":"2024-04-06T07:29:13.084746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = df.merge(pdf[['patient_id','patient_seizure','patient_lpd','patient_gpd','patient_lrda','patient_grda','patient_other']])","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:29:13.087149Z","iopub.execute_input":"2024-04-06T07:29:13.08742Z","iopub.status.idle":"2024-04-06T07:29:13.099782Z","shell.execute_reply.started":"2024-04-06T07:29:13.087398Z","shell.execute_reply":"2024-04-06T07:29:13.098867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# df['path'] = '/kaggle/input/create-fixed-eeg-spec/fmax30_sec30_8ims/' + df.eeg_id.astype(str) + '_' + df.group_in_eeg.astype(str) + '.npy'\n# df.to_csv('/kaggle/working/train.csv', index=False)\n# for fold in [0]:\n#     !python3 train_one_fold.py -c finetune_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg -e 1 --fold {fold}\n# for fold in [3,4]:\n#     !python3 train_one_fold.py -c finetune_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg -e 7 --fold {fold}\n\n\n\n# df['path'] = '/kaggle/input/fmax90-sec10-8ims/fmax90_sec10_8ims/' + df.eeg_id.astype(str) + '_' + df.group_in_eeg.astype(str) + '.npy'\n# df.to_csv('/kaggle/working/train.csv', index=False)\n# for fold in range(5):\n#     !python3 train_one_fold.py -c finetune_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg -e 7 --fold {fold}\n\n\n\ndf['path'] = '/kaggle/input/fmax60-sec40-16ims/fmax60_sec40_16ims/' + df.eeg_id.astype(str) + '_' + df.group_in_eeg.astype(str) + '.npy'\ndf.to_csv('/kaggle/working/train.csv', index=False)\nfor fold in [1,2]:\n    !python3 train_one_fold.py -c finetune_hms_chris_fmax60_40sec_16ims_bandpass -e 9 --fold {fold}\n# for fold in [3,4]:\n#     !python3 train_one_fold.py -c finetune_hms_chris_fmax60_40sec_16ims_bandpass -e 9 --fold {fold}    \n\n# df['d'] = 'v1-fmax60-sec50-win256-16ims'\n# df.loc[df.eeg_id>=2104387418, 'd'] = 'v2-fmax60-sec50-win256-16ims'\n# df['path'] = '/kaggle/input/'+ df.d + '/fmax60_sec50_win256_16ims/' + df.eeg_id.astype(str) + '_' + df.group_in_eeg.astype(str) + '.npy'\n# df.to_csv('/kaggle/working/train.csv', index=False)\n# for fold in [0]:\n#     !python3 train_one_fold.py -c finetune_hms_chris_fmax60_50sec_win256_16ims_bandpass_spe_and_eeg -e 14 --fold {fold}\n# for fold in [3,4]:\n#     !python3 train_one_fold.py -c finetune_hms_chris_fmax60_50sec_win256_16ims_bandpass_spe_and_eeg -e 14 --fold {fold}    ","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:29:13.101541Z","iopub.execute_input":"2024-04-06T07:29:13.101958Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import os\n# for path in df.path:\n#     assert os.path.exists(path)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:28:54.419014Z","iopub.execute_input":"2024-04-06T07:28:54.419362Z","iopub.status.idle":"2024-04-06T07:28:54.423914Z","shell.execute_reply.started":"2024-04-06T07:28:54.419332Z","shell.execute_reply":"2024-04-06T07:28:54.42285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !ls /kaggle/input/v2-fmax60-sec50-win256-16ims/fmax60_sec50_win256_16ims/2104861236_0.npy\n# !ls /kaggle/input/v2-fmax60-sec50-win256-16ims/fmax60_sec50_win256_16ims/3484900135*","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:28:54.425095Z","iopub.execute_input":"2024-04-06T07:28:54.425371Z","iopub.status.idle":"2024-04-06T07:28:54.434154Z","shell.execute_reply.started":"2024-04-06T07:28:54.425348Z","shell.execute_reply":"2024-04-06T07:28:54.433187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# df.sample(40).to_csv('/kaggle/working/train.csv', index=False)\n","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:28:54.436214Z","iopub.execute_input":"2024-04-06T07:28:54.436512Z","iopub.status.idle":"2024-04-06T07:28:54.447478Z","shell.execute_reply.started":"2024-04-06T07:28:54.436487Z","shell.execute_reply":"2024-04-06T07:28:54.44661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !python3 train_one_fold.py -c finetune_hms_chris_fmax60_40sec_16ims_bandpass -e 1 --fold 1","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:28:54.44891Z","iopub.execute_input":"2024-04-06T07:28:54.449565Z","iopub.status.idle":"2024-04-06T07:28:54.456492Z","shell.execute_reply.started":"2024-04-06T07:28:54.449529Z","shell.execute_reply":"2024-04-06T07:28:54.454894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# val metric: tensor(-0.2644) # 1ep\n","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:28:54.458027Z","iopub.execute_input":"2024-04-06T07:28:54.458838Z","iopub.status.idle":"2024-04-06T07:28:54.46564Z","shell.execute_reply.started":"2024-04-06T07:28:54.458782Z","shell.execute_reply":"2024-04-06T07:28:54.464741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%cd /kaggle/working","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:28:54.466769Z","iopub.execute_input":"2024-04-06T07:28:54.467054Z","iopub.status.idle":"2024-04-06T07:28:54.476864Z","shell.execute_reply.started":"2024-04-06T07:28:54.46703Z","shell.execute_reply":"2024-04-06T07:28:54.475854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import albumentations as A\n\ndef hms_preprocess_concat(cfg, row):\n    spe_path = row.spe_path\n    image = pd.read_parquet(spe_path)\n    if debug:\n        t = row.spectrogram_label_offset_seconds\n        image = image[(image.time>t) & (image.time<t+600)]\n    image = image.values[:, 1:]\n        \n    if cfg.use_center_sec is not None:\n        cut = (300-cfg.use_center_sec)//2\n        image = image[cut:-cut, :]\n        assert image.shape[0] == cfg.use_center_sec\n\n    image = np.clip(image, np.exp(-4), np.exp(8))\n    image = np.log(image)\n\n    ep = 1e-6\n    mu, std = (-0.13223082279419138, 2.368639211554277)\n    image = (image-mu)/(std+ep)\n    image = np.nan_to_num(image, nan=0.0)\n    image = np.array([image, image, image]).transpose((1,2,0))\n    return image\n\nclass ClassificationDataset(Dataset):\n    def __init__(self, df, cfg):\n        self.df = df\n        self.cfg = cfg\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image = hms_preprocess_concat(self.cfg, row)\n\n        image = self.cfg.transform(image=image)['image']\n        image = torch.FloatTensor(image.transpose((2, 0, 1)))\n\n        return image\n\ndef load_eeg_spectrograms_concat(cfg, path):\n    image = np.load(path)\n    image = (image-cfg.mu_std[0])/cfg.mu_std[1]\n    \n    image = np.concatenate([image[:,:,i] for i in range(image.shape[2])], 0)\n#     print(image.shape)\n    image = np.array([image, image, image]).transpose((1,2,0))\n    \n    return image\n\nclass ClassificationDatasetEeg(Dataset):\n    def __init__(self, df, cfg):\n        self.paths = df.path.values\n        self.cfg = cfg\n\n    def __len__(self):\n        return len(self.paths)\n\n    def __getitem__(self, idx):\n        path = self.paths[idx]\n        image = load_eeg_spectrograms_concat(self.cfg, path)\n\n        image = self.cfg.transform(image=image)['image']\n        image = torch.FloatTensor(image.transpose((2, 0, 1)))\n\n        return image    \n    \nclass ClassificationDatasetSpeEegConcat(Dataset):    \n    def __init__(self, df, cfg):\n        self.df = df\n        self.cfg = cfg\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image = hms_preprocess_concat(self.cfg, row)\n        add_im = load_eeg_spectrograms_concat(self.cfg, row.path)\n        if not ((hasattr(self.cfg, 'skip_add_im_resize')) and (self.cfg.skip_add_im_resize)):\n            add_im = cv2.resize(add_im, (512, add_im.shape[1])).transpose((1,0,2))\n        image = cv2.resize(image, (512, 256)).transpose((1,0,2))\n        \n        image = np.concatenate([image, add_im], axis=1)\n\n        image = self.cfg.transform(image=image)['image']\n        image = torch.FloatTensor(image.transpose((2, 0, 1)))\n\n        return image","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:28:54.773469Z","iopub.execute_input":"2024-04-06T07:28:54.773911Z","iopub.status.idle":"2024-04-06T07:28:54.795776Z","shell.execute_reply.started":"2024-04-06T07:28:54.77386Z","shell.execute_reply":"2024-04-06T07:28:54.794784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class hms_base():\n    def __init__(self):\n        super().__init__()\n        self.label_features = ['seizure','lpd','gpd','lrda','grda','other']\n        self.num_classes = len(self.label_features)\n        self.model_name = 'swinv2_large_window12to24_192to384.ms_in22k_ft_in1k'\n        self.model = timm.create_model(self.model_name, pretrained=False, num_classes=self.num_classes)\n        self.image_size = (384, 384)\n        self.transform = A.Compose([A.Resize(self.image_size[0], self.image_size[1])])\n        self.batch_size = 2\n        self.use_center_sec = None\n        self.tta = None","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:28:54.962638Z","iopub.execute_input":"2024-04-06T07:28:54.962939Z","iopub.status.idle":"2024-04-06T07:28:54.969348Z","shell.execute_reply.started":"2024-04-06T07:28:54.962915Z","shell.execute_reply":"2024-04-06T07:28:54.968472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict_one_fmax(config_list, alldf, fold=0):\n    df = alldf[alldf.fold==fold]\n    all_preds = []\n    for i, cfg_name in enumerate(config_list):\n        cfg = eval(cfg_name)()\n        cfg.df = copy.deepcopy(df)\n        models = []\n        for model_path in cfg.model_paths:\n            if f'last_fold{fold}' not in model_path:\n                continue\n            print(model_path)\n            model = copy.deepcopy(cfg.model)\n            torch_state_dict = torch.load(model_path, map_location=torch.device('cpu'))\n            if 'state_dict' in torch_state_dict:\n                torch_state_dict = torch_state_dict['state_dict']\n                fix_state_dict = {}\n                for k, v in torch_state_dict.items():\n                    fix_state_dict[k[6:]] = v\n                model.load_state_dict(fix_state_dict)\n            else:\n                model.load_state_dict(torch_state_dict)\n            model.to(DEVICE)\n            model.eval()\n            models.append(model)\n        if (hasattr(cfg, 'kaggle_spe')) and (cfg.kaggle_spe):\n            ds = ClassificationDataset(df, cfg)\n        elif (hasattr(cfg, 'spe_eeg_concat')) and (cfg.spe_eeg_concat):\n            print('spe_eeg_concat')\n            ds = ClassificationDatasetSpeEegConcat(df, cfg)\n        else:\n            ds = ClassificationDatasetEeg(df, cfg)\n        loader = DataLoader(ds, batch_size=cfg.batch_size, shuffle=False, drop_last=False, num_workers=2)\n        preds = []\n        for images in tqdm(loader, smoothing=0):\n            images = images.to(DEVICE)\n            batch_preds = []\n            for model in models:\n                batch_preds.append(model(images).detach().cpu().numpy())\n\n            if cfg.tta:\n                imsize = images.size()\n                images = torch.flip(images, (3,))\n                assert imsize == images.size()\n                for model in models:\n                    batch_preds.append(model(images).detach().cpu().numpy())\n\n            preds += np.mean(batch_preds, axis=0).tolist()\n        fs = [f'{cfg_name}_pred_{col}' for col in cfg.label_features]\n        cfg.df[fs] = preds\n        cfg.df.to_csv(f'pred_{cfg_name}_fold{fold}.csv', index=False)\n        all_preds.append(preds)\n        if debug:\n            oof = df.copy()\n            true_cols = ['seizure','lpd','gpd','lrda','grda','other']\n            pred_cols = [f'pred_{c}' for c in true_cols]\n            oof[pred_cols] = softmax(preds, 1)\n            true = oof[true_cols]\n            pred = oof[pred_cols]\n            true['id'] = list(range(len(true)))\n            pred['id'] = list(range(len(pred)))\n            true.columns = [0,1,2,3,4,5,'id']\n            pred.columns = [0,1,2,3,4,5,'id']\n            score = kaggle_metric(solution=true, submission=pred, row_id_column_name='id')\n            print(cfg_name, round(score, 4))\n    #     break\n#     if debug:\n#         oof = df.copy()\n#         true_cols = ['seizure','lpd','gpd','lrda','grda','other']\n#         pred_cols = [f'pred_{c}' for c in true_cols]\n#         oof[pred_cols] = softmax(np.mean(all_preds, 0), 1)\n#         true = oof[true_cols]\n#         pred = oof[pred_cols]\n#         true['id'] = list(range(len(true)))\n#         pred['id'] = list(range(len(pred)))\n#         true.columns = [0,1,2,3,4,5,'id']\n#         pred.columns = [0,1,2,3,4,5,'id']\n#         score = kaggle_metric(solution=true, submission=pred, row_id_column_name='id')\n#         print(round(score, 4))","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:27:45.475008Z","iopub.execute_input":"2024-04-06T07:27:45.47525Z","iopub.status.idle":"2024-04-06T07:27:45.49739Z","shell.execute_reply.started":"2024-04-06T07:27:45.475229Z","shell.execute_reply":"2024-04-06T07:27:45.496313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# fmax30, 30sec, bandpass\nclass finetune_hms_chris_fmax30_30sec_8ims_bandpass(hms_base):\n    def __init__(self):\n        super().__init__()\n        self.model_name = 'swinv2_large_window12to24_192to384.ms_in22k_ft_in1k'\n        self.model = timm.create_model(self.model_name, pretrained=False, num_classes=self.num_classes)\n        self.model_paths = [\n            '/kaggle/input/train-by-fixed-data-e24ea3/hms/results/finetune_hms_chris_fmax30_30sec_8ims_bandpass/last_fold0.ckpt',\n            '/kaggle/input/train-by-fixed-data-e24ea3/hms/results/finetune_hms_chris_fmax30_30sec_8ims_bandpass/last_fold1.ckpt',\n            '/kaggle/input/train-by-fixed-data-e24ea3/hms/results/finetune_hms_chris_fmax30_30sec_8ims_bandpass/last_fold2.ckpt',\n            '/kaggle/input/train-by-fixed-data/hms/results/finetune_hms_chris_fmax30_30sec_8ims_bandpass/last_fold3.ckpt',\n            '/kaggle/input/train-by-fixed-data/hms/results/finetune_hms_chris_fmax30_30sec_8ims_bandpass/last_fold4.ckpt'\n        ]\n        self.image_size = (384, 384)\n        self.transform = A.Compose([A.Resize(self.image_size[0], self.image_size[1])])\n        self.mu_std = (9, 3.5)\n\nclass finetune_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg(hms_base):\n    def __init__(self):\n        super().__init__()\n        self.model_name = 'swinv2_large_window12to24_192to384.ms_in22k_ft_in1k'\n        self.model = timm.create_model(self.model_name, pretrained=False, num_classes=self.num_classes)\n        self.model_paths = [f'/kaggle/working/hms/results/finetune_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg/last_fold{train_fold}.ckpt']\n        self.image_size = (384, 384)\n        self.transform = A.Compose([A.Resize(self.image_size[0], self.image_size[1])])\n        self.mu_std = (9, 3.5)\n        self.spe_eeg_concat = True\n\nconfig_list = [\n#     'finetune_hms_chris_fmax30_30sec_8ims_bandpass',\n    'finetune_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg'\n]\n# predict_one_fmax(config_list, df, fold=0)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:27:45.498624Z","iopub.execute_input":"2024-04-06T07:27:45.498961Z","iopub.status.idle":"2024-04-06T07:27:45.512567Z","shell.execute_reply.started":"2024-04-06T07:27:45.498934Z","shell.execute_reply":"2024-04-06T07:27:45.511831Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for fold in range(5):\n#     predict_one_fmax(config_list, df, fold=fold)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:27:45.513805Z","iopub.execute_input":"2024-04-06T07:27:45.514176Z","iopub.status.idle":"2024-04-06T07:27:45.524425Z","shell.execute_reply.started":"2024-04-06T07:27:45.514125Z","shell.execute_reply":"2024-04-06T07:27:45.523518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# oof=pd.concat([pd.read_csv(f'/kaggle/working/pred_finetune_hms_chris_fmax30_30sec_8ims_bandpass_fold{fold}.csv') for fold in range(5)])\n# oof.to_csv('/kaggle/working/pred_finetune_hms_chris_fmax30_30sec_8ims_bandpass.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T07:27:45.525565Z","iopub.execute_input":"2024-04-06T07:27:45.525918Z","iopub.status.idle":"2024-04-06T07:27:45.536885Z","shell.execute_reply.started":"2024-04-06T07:27:45.525893Z","shell.execute_reply":"2024-04-06T07:27:45.53606Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}