{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","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":7979206,"sourceType":"datasetVersion","datasetId":4596419},{"sourceId":7981321,"sourceType":"datasetVersion","datasetId":4419986},{"sourceId":7989915,"sourceType":"datasetVersion","datasetId":4703503},{"sourceId":8015225,"sourceType":"datasetVersion","datasetId":4435105},{"sourceId":8016808,"sourceType":"datasetVersion","datasetId":4723312},{"sourceId":8025136,"sourceType":"datasetVersion","datasetId":4729369},{"sourceId":8035920,"sourceType":"datasetVersion","datasetId":4737176},{"sourceId":8054458,"sourceType":"datasetVersion","datasetId":4724042},{"sourceId":8055494,"sourceType":"datasetVersion","datasetId":4687290},{"sourceId":8077221,"sourceType":"datasetVersion","datasetId":4708755},{"sourceId":8105078,"sourceType":"datasetVersion","datasetId":4701635},{"sourceId":8109988,"sourceType":"datasetVersion","datasetId":4652811},{"sourceId":162276846,"sourceType":"kernelVersion"},{"sourceId":170279329,"sourceType":"kernelVersion"},{"sourceId":170279396,"sourceType":"kernelVersion"},{"sourceId":170388873,"sourceType":"kernelVersion"},{"sourceId":170625834,"sourceType":"kernelVersion"},{"sourceId":170714546,"sourceType":"kernelVersion"},{"sourceId":170714548,"sourceType":"kernelVersion"},{"sourceId":170715357,"sourceType":"kernelVersion"},{"sourceId":170739282,"sourceType":"kernelVersion"},{"sourceId":170742054,"sourceType":"kernelVersion"},{"sourceId":170742073,"sourceType":"kernelVersion"},{"sourceId":170763234,"sourceType":"kernelVersion"},{"sourceId":170763274,"sourceType":"kernelVersion"},{"sourceId":170787277,"sourceType":"kernelVersion"},{"sourceId":170806449,"sourceType":"kernelVersion"},{"sourceId":170882149,"sourceType":"kernelVersion"},{"sourceId":170883254,"sourceType":"kernelVersion"},{"sourceId":170883282,"sourceType":"kernelVersion"},{"sourceId":170883536,"sourceType":"kernelVersion"},{"sourceId":170883539,"sourceType":"kernelVersion"},{"sourceId":170912829,"sourceType":"kernelVersion"},{"sourceId":171749895,"sourceType":"kernelVersion"}],"dockerImageVersionId":30636,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# tattaka","metadata":{}},{"cell_type":"code","source":"# import pandas as pd\n# if len(pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/sample_submission.csv'))==1:\n#     pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/sample_submission.csv').to_csv('submission.csv', index=False)\n#     raise","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:51:00.077268Z","iopub.execute_input":"2024-04-14T05:51:00.077977Z","iopub.status.idle":"2024-04-14T05:51:00.082463Z","shell.execute_reply.started":"2024-04-14T05:51:00.077950Z","shell.execute_reply":"2024-04-14T05:51:00.081466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile tattaka.py\nfrom glob import glob\n\nimport numpy as np\nimport pandas as pd\nimport pytorch_lightning as pl\nimport torch\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom torch.nn import functional as F\nfrom tqdm.auto import tqdm\n\nimport numpy as np\nimport pandas as pd\nimport pytorch_lightning as pl\nimport timm\nimport torch\nfrom pytorch_lightning import LightningDataModule, callbacks\nfrom pytorch_lightning.loggers import WandbLogger\nfrom pytorch_lightning.utilities import rank_zero_info\nfrom scipy.signal import butter, lfilter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom timm.utils import ModelEmaV2\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom torch.nn.parameter import Parameter\nfrom torch.optim import AdamW\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchaudio.transforms import AmplitudeToDB, MelSpectrogram\nfrom transformers import get_cosine_schedule_with_warmup\nfrom typing import List\nimport librosa\n\nimport sys\nmode = \"test\"\nbatch_size = 4\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nBASE_DIR = \"/kaggle/input/hms-harmful-brain-activity-classification/\"\nEEG_PATH = f\"{BASE_DIR}/{mode}_eegs/\"\nSPEC_PATH = f\"{BASE_DIR}/{mode}_spectrograms/\"\n\n\n# sys.path.append(\"/kaggle/input/hms-weights-2\")\n# from main_exp092 import HMSLightningModel, HMSDataset\n\n# exp_dirs = [\n#     \"/kaggle/input/hms-weights-2/exp108/convnext_large_384_el40_mixup_50ep\",\n# ]\n# # exp_weights = np.ones(len(exp_dirs)) / len(exp_dirs)\n# exp_weights = [1]\n# for exp_dir in exp_dirs:\n#     try:\n#         df = pd.read_csv(glob(f\"{exp_dir}/**/result_df.csv\", recursive=True)[0])\n#         print(f\"{exp_dir}: \", df.kl_div.mean())\n#     except:\n#         pass\n\n# pl.seed_everything(42, workers=True)\n# df = pd.read_csv(f\"/kaggle/input/hms-harmful-brain-activity-classification/{mode}.csv\")\n\n# preds_all = []\n# eeg_ids_all = []\n# for exp_dir in exp_dirs:\n#     config = exp_dir.split('/')[-1]\n#     model_paths = glob(f\"{exp_dir}/**/best_loss.ckpt\", recursive=True)\n#     print(model_paths)\n#     models = [\n#         HMSLightningModel.load_from_checkpoint(model_path, pretrained=False,).eval().to(device=device)\n#         for model_path in model_paths\n#     ]\n#     dataset = HMSDataset(\n#         df=df,\n#         overlap_df=None,\n#         fmin = models[0].fmin,\n#         fmax = models[0].fmax,\n#         mode=\"test\",\n#         eeg_path=EEG_PATH,\n#         spec_path=SPEC_PATH,\n#     )\n#     dataloader = DataLoader(\n#         dataset=dataset,\n#         batch_size=batch_size,\n#         num_workers=2,\n#         shuffle=False,\n#         drop_last=False,\n#         pin_memory=True,\n#     )\n#     preds_exp = []\n#     eeg_ids_all = []\n#     for batch in tqdm(dataloader):\n#         eeg_ids, signals, specs = batch[\"eeg_id\"], batch[\"signals\"], batch[\"specs\"]\n#         eeg_ids_all.append(eeg_ids.numpy())\n#         image = models[0].pipeline(signals.to(device=device))\n#         specs = specs.to(device=device)\n#         with torch.no_grad():\n#             preds_exp.append(np.stack([model.model_ema.module(image, specs).detach().cpu().numpy() for model in models]).mean(0))\n#     preds_exp = np.concatenate(preds_exp)\n#     eeg_ids_all = np.concatenate(eeg_ids_all)\n\n#     submission = pd.DataFrame({\"eeg_id\": eeg_ids_all})\n#     submission[[\"seizure_vote\", \"lpd_vote\", \"gpd_vote\", \"lrda_vote\", \"grda_vote\", \"other_vote\"]] = preds_exp\n#     submission.to_csv(f\"{config}.csv\",index=False)\n#     submission.head()\n\n\nsys.path.append(\"/kaggle/input/hms-weights-2\")\nfrom main_exp094 import HMSLightningModel2_5D, HMSDataset\n\nexp_dirs = [\n    \"/kaggle/input/hms-weights-2/exp094/caformer_s18_2_5d_256_el30_mixup_100ep\",\n]\n\nfor exp_dir in exp_dirs:\n    try:\n        df = pd.read_csv(glob(f\"{exp_dir}/**/result_df.csv\", recursive=True)[0])\n        print(f\"{exp_dir}: \", df.kl_div.mean())\n    except:\n        pass\n\npl.seed_everything(42, workers=True)\ndf = pd.read_csv(f\"/kaggle/input/hms-harmful-brain-activity-classification/{mode}.csv\")\n\npreds_all = []\neeg_ids_all = []\nfor exp_dir in exp_dirs:\n    config = exp_dir.split('/')[-1]\n    model_paths = glob(f\"{exp_dir}/**/best_loss.ckpt\", recursive=True)\n    print(model_paths)\n    models = [\n        HMSLightningModel2_5D.load_from_checkpoint(model_path, pretrained=False,).eval().to(device=device)\n        for model_path in model_paths\n    ]\n    dataset = HMSDataset(\n        df=df,\n        overlap_df=None,\n        fmin = models[0].fmin,\n        fmax = models[0].fmax,\n        mode=\"test\",\n        eeg_path=EEG_PATH,\n        spec_path=SPEC_PATH,\n    )\n    dataloader = DataLoader(\n        dataset=dataset,\n        batch_size=batch_size,\n        num_workers=2,\n        shuffle=False,\n        drop_last=False,\n        pin_memory=True,\n    )\n    preds_exp = []\n    eeg_ids_all = []\n    for batch in tqdm(dataloader):\n        eeg_ids, signals, specs = batch[\"eeg_id\"], batch[\"signals\"], batch[\"specs\"]\n        eeg_ids_all.append(eeg_ids.numpy())\n        image = models[0].pipeline(signals.to(device=device))\n        specs = specs.to(device=device)\n        with torch.no_grad():\n            preds_exp.append(np.stack([model.model_ema.module(image, specs).detach().cpu().numpy() for model in models]).mean(0))\n    preds_exp = np.concatenate(preds_exp)\n    eeg_ids_all = np.concatenate(eeg_ids_all)\n\n    submission = pd.DataFrame({\"eeg_id\": eeg_ids_all})\n    submission[[\"seizure_vote\", \"lpd_vote\", \"gpd_vote\", \"lrda_vote\", \"grda_vote\", \"other_vote\"]] = preds_exp\n    submission.to_csv(f\"{config}.csv\",index=False)\n    submission.head()\n\nsys.path.append(\"/kaggle/input/hms-weights-3\")\nfrom main_exp147 import HMSLightningModel, HMSDataset \n\nexp_dirs = [\n    \"/kaggle/input/hms-weights-3/exp147/tiny_vit_21m_512_el30_mixup_50ep\",\n]\n\nfor exp_dir in exp_dirs:\n    try:\n        df = pd.read_csv(glob(f\"{exp_dir}/**/result_df.csv\", recursive=True)[0])\n        print(f\"{exp_dir}: \", df.kl_div.mean())\n    except:\n        pass\n\npl.seed_everything(42, workers=True)\ndf = pd.read_csv(f\"/kaggle/input/hms-harmful-brain-activity-classification/{mode}.csv\")\n\npreds_all = []\neeg_ids_all = []\nfor exp_dir in exp_dirs:\n    config = exp_dir.split('/')[-1]\n    model_paths = glob(f\"{exp_dir}/**/best_loss.ckpt\", recursive=True)\n    print(model_paths)\n    models = [\n        HMSLightningModel.load_from_checkpoint(model_path, pretrained=False,).eval().to(device=device)\n        for model_path in model_paths\n    ]\n    dataset = HMSDataset(\n        df=df,\n        overlap_df=None,\n        fmin = models[0].fmin,\n        fmax = models[0].fmax,\n        mode=\"test\",\n        eeg_path=EEG_PATH,\n        spec_path=SPEC_PATH,\n    )\n    dataloader = DataLoader(\n        dataset=dataset,\n        batch_size=batch_size,\n        num_workers=2,\n        shuffle=False,\n        drop_last=False,\n        pin_memory=True,\n    )\n    preds_exp = []\n    eeg_ids_all = []\n    for batch in tqdm(dataloader):\n        eeg_ids, signals, specs = batch[\"eeg_id\"], batch[\"signals\"], batch[\"specs\"]\n        eeg_ids_all.append(eeg_ids.numpy())\n        image = models[0].pipeline(signals.to(device=device))\n        specs = specs.to(device=device)\n        with torch.no_grad():\n            preds_exp.append(np.stack([model.model_ema.module(image, specs).detach().cpu().numpy() for model in models]).mean(0))\n    preds_exp = np.concatenate(preds_exp)\n    eeg_ids_all = np.concatenate(eeg_ids_all)\n\n    submission = pd.DataFrame({\"eeg_id\": eeg_ids_all})\n    submission[[\"seizure_vote\", \"lpd_vote\", \"gpd_vote\", \"lrda_vote\", \"grda_vote\", \"other_vote\"]] = preds_exp\n    submission.to_csv(f\"{config}.csv\",index=False)\n    submission.head()\n    ","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:51:00.308542Z","iopub.execute_input":"2024-04-14T05:51:00.308863Z","iopub.status.idle":"2024-04-14T05:51:00.325964Z","shell.execute_reply.started":"2024-04-14T05:51:00.308838Z","shell.execute_reply":"2024-04-14T05:51:00.325050Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!python tattaka.py","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:51:00.327423Z","iopub.execute_input":"2024-04-14T05:51:00.327992Z","iopub.status.idle":"2024-04-14T05:52:01.782921Z","shell.execute_reply.started":"2024-04-14T05:51:00.327966Z","shell.execute_reply":"2024-04-14T05:52:01.781776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# bilzard","metadata":{}},{"cell_type":"code","source":"# %%writefile bilzard.py","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:52:01.784214Z","iopub.execute_input":"2024-04-14T05:52:01.784533Z","iopub.status.idle":"2024-04-14T05:52:01.788996Z","shell.execute_reply.started":"2024-04-14T05:52:01.784504Z","shell.execute_reply":"2024-04-14T05:52:01.787998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def is_interactive():\n    return \"runtime\" in get_ipython().config.IPKernelApp.connection_file\n\nprint(\"is_interactive\", is_interactive())","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:52:01.790938Z","iopub.execute_input":"2024-04-14T05:52:01.791259Z","iopub.status.idle":"2024-04-14T05:52:01.802134Z","shell.execute_reply.started":"2024-04-14T05:52:01.791225Z","shell.execute_reply":"2024-04-14T05:52:01.801309Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ENSEMBLE_ENTITY = \"bilzard_v1\"\nENSEMBLE_ENTITY_NAME = \"xxx\"\nDEBUG = True","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:52:01.803195Z","iopub.execute_input":"2024-04-14T05:52:01.803535Z","iopub.status.idle":"2024-04-14T05:52:01.815923Z","shell.execute_reply.started":"2024-04-14T05:52:01.803502Z","shell.execute_reply":"2024-04-14T05:52:01.815147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not is_interactive():\n    DEBUG = False","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:52:01.816867Z","iopub.execute_input":"2024-04-14T05:52:01.817109Z","iopub.status.idle":"2024-04-14T05:52:01.825719Z","shell.execute_reply.started":"2024-04-14T05:52:01.817087Z","shell.execute_reply":"2024-04-14T05:52:01.824963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"DEBUG:\", DEBUG)\nprint(\"ENSEMBLE_ENTITY\", ENSEMBLE_ENTITY)\nprint(\"ENSEMBLE_ENTITY_NAME\", ENSEMBLE_ENTITY_NAME)","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:52:01.826806Z","iopub.execute_input":"2024-04-14T05:52:01.827575Z","iopub.status.idle":"2024-04-14T05:52:01.839469Z","shell.execute_reply.started":"2024-04-14T05:52:01.827543Z","shell.execute_reply":"2024-04-14T05:52:01.838571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%cd /kaggle/working","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:52:01.840665Z","iopub.execute_input":"2024-04-14T05:52:01.840985Z","iopub.status.idle":"2024-04-14T05:52:01.850764Z","shell.execute_reply.started":"2024-04-14T05:52:01.840955Z","shell.execute_reply":"2024-04-14T05:52:01.849918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -U hydra-core polars einops \\\n    --no-index \\\n    --find-links=/kaggle/input/hms-build-runtime-environment/wheels","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:52:01.854659Z","iopub.execute_input":"2024-04-14T05:52:01.854981Z","iopub.status.idle":"2024-04-14T05:52:28.542968Z","shell.execute_reply.started":"2024-04-14T05:52:01.854958Z","shell.execute_reply":"2024-04-14T05:52:28.541998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! rm -rf /kaggle/temp && \\\n    cp -r /kaggle/input/hms-clone-kagle-hms-bilzard-repo /kaggle/temp","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:52:28.544286Z","iopub.execute_input":"2024-04-14T05:52:28.544586Z","iopub.status.idle":"2024-04-14T05:52:33.513575Z","shell.execute_reply.started":"2024-04-14T05:52:28.544559Z","shell.execute_reply":"2024-04-14T05:52:33.512201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%cd /kaggle/temp/kaggle-hms-bilzard","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:52:33.515296Z","iopub.execute_input":"2024-04-14T05:52:33.516091Z","iopub.status.idle":"2024-04-14T05:52:33.522675Z","shell.execute_reply.started":"2024-04-14T05:52:33.516050Z","shell.execute_reply":"2024-04-14T05:52:33.521763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!git log --oneline -1","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:52:33.523928Z","iopub.execute_input":"2024-04-14T05:52:33.524195Z","iopub.status.idle":"2024-04-14T05:52:34.507692Z","shell.execute_reply.started":"2024-04-14T05:52:33.524172Z","shell.execute_reply":"2024-04-14T05:52:34.506780Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%writefile run/new_batch_infer.py","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:52:34.509021Z","iopub.execute_input":"2024-04-14T05:52:34.509335Z","iopub.status.idle":"2024-04-14T05:52:34.513725Z","shell.execute_reply.started":"2024-04-14T05:52:34.509307Z","shell.execute_reply":"2024-04-14T05:52:34.512773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile run/new_batch_infer.py\nfrom pathlib import Path\nfrom typing import cast\n\nimport hydra\nimport numpy as np\nimport polars as pl\nimport torch.nn as nn\nfrom hydra.core.global_hydra import GlobalHydra\nfrom torch.utils.data import DataLoader\n\nfrom src.config import EnsembleExperimentConfig, EnsembleMainConfig, MainConfig\nfrom src.constant import LABELS\nfrom src.data_util import preload_cqf, preload_eegs, preload_spectrograms\nfrom src.evaluator import Evaluator\nfrom src.infer_util import load_metadata, make_submission\nfrom src.logger import BaseLogger\nfrom src.proc_util import trace\nfrom src.random_util import seed_everything\nfrom src.train_util import check_model, get_model\n\nfrom .ensemble import do_evaluate\nfrom .infer import get_loader, load_checkpoint, predict\n\n\ndef load_config(config_name, parent_cfg: EnsembleMainConfig) -> MainConfig:\n    if GlobalHydra.instance().is_initialized():\n        GlobalHydra.instance().clear()\n\n    with hydra.initialize(config_path=\"conf\", version_base=\"1.2\"):\n        cfg = hydra.compose(\n            config_name=config_name,\n            overrides=[\n                f\"phase={parent_cfg.phase}\",\n                f\"env={parent_cfg.env.name}\",\n                f\"infer.batch_size={parent_cfg.env.infer_batch_size}\",\n                f\"architecture.model.encoder.grad_checkpointing={parent_cfg.env.grad_checkpointing}\",\n            ],\n        )\n        print(\"** phase:\", cfg.phase)\n        print(\"** env:\", cfg.env.name)\n        print(\"** infer_batch_size:\", cfg.infer.batch_size)\n        print(\n            \"** grad_checkpointing:\", cfg.architecture.model.encoder.grad_checkpointing\n        )\n\n    return cast(MainConfig, cfg)\n\n\ndef infer_per_seed(\n    cfg: MainConfig, model: nn.Module, data_loader: DataLoader, seed: int\n) -> pl.DataFrame:\n    # TODO: 各fold/seedごとの推論結果をparquetでinferディレクトリ配下に保存する\n    seed_everything(seed)\n\n    match cfg.phase:\n        # TODO: evaluateの有無をオプションで指定できるようにする\n        case \"train\":\n            evaluator = Evaluator(\n                aggregation_fn=cfg.trainer.val.aggregation_fn,\n                input_keys=cfg.trainer.data.input_keys,\n                agg_policy=cfg.trainer.val.agg_policy,\n                iterations=cfg.infer.tta_iterations,\n                weight_exponent=cfg.trainer.val.weight_exponent,\n                min_weight=cfg.trainer.val.min_weight,\n            )\n            output = evaluator.evaluate(model, data_loader)\n            val_loss, val_loss_per_label, eeg_ids, logits = (\n                output[\"val_loss\"],\n                output[\"val_loss_per_label\"],\n                output[\"eeg_ids\"],\n                output[\"logits_per_eeg\"],\n            )\n            print(f\"val_loss: {val_loss:.4f}\")\n            print(\", \".join([f\"{k}={v:.4f}\" for k, v in val_loss_per_label.items()]))\n        case \"test\" | \"develop\":\n            eeg_ids, logits = predict(\n                model,\n                data_loader,\n                cfg.trainer.data.input_keys,\n                iterations=cfg.infer.tta_iterations,\n            )\n        case _:\n            raise ValueError(f\"Invalid phase: {cfg.phase}\")\n\n    prediction_df = make_submission(eeg_ids, logits, apply_softmax=False)\n    return prediction_df\n\n\ndef infer_per_experiment(\n    parent_config: EnsembleMainConfig,\n    experiment: EnsembleExperimentConfig,\n    metadata: pl.DataFrame,\n    id2eeg: dict[int, np.ndarray],\n    id2cqf: dict[int, np.ndarray],\n    spec_id2spec: dict[int, np.ndarray],\n) -> pl.DataFrame:\n    \"\"\"\n    process per experiment\n    1. load config\n    2. load data loader\n    3. for each fold and seed:\n        - load model weight and do inference\n    \"\"\"\n    cfg = load_config(experiment.exp_name, parent_config)\n    logger = BaseLogger(log_file_name=cfg.infer.log_name, clear=True)\n    model = get_model(cfg.architecture, pretrained=False)\n    check_model(cfg.architecture, model)\n\n    check_loader = get_loader(\n        cfg=cfg,\n        metadata=metadata.sample(1),\n        id2eeg=id2eeg,\n        id2cqf=id2cqf,\n        spec_id2spec=spec_id2spec,\n    )\n    logger.write_log(\"Dataset:\", check_loader.dataset)\n    logger.write_log(\"Model:\", model)\n\n    del check_loader\n\n    metadata_all = metadata\n    pred_dfs = []\n    for ensemble_fold in experiment.folds:\n        fold = ensemble_fold.split\n        # trainの場合、validation dataのみ推論する\n        if cfg.phase == \"train\":\n            metadata = metadata_all.filter(pl.col(\"fold\").eq(fold))\n\n        data_loader = get_loader(\n            cfg=cfg,\n            metadata=metadata,\n            id2eeg=id2eeg,\n            id2cqf=id2cqf,\n            spec_id2spec=spec_id2spec,\n        )\n\n        for seed in ensemble_fold.seeds:\n            print(\"*\" * 50)\n            print(f\"* exp_name: {cfg.exp_name}, fold: {fold}, seed: {seed}\")\n            print(\"*\" * 50)\n            weight_path = (\n                Path(cfg.env.checkpoint_dir)\n                / cfg.exp_name\n                / f\"fold_{fold}\"\n                / f\"seed_{seed}\"\n                / \"model\"\n                / f\"{cfg.infer.model_choice}_model.pth\"\n            )\n            load_checkpoint(model, weight_path)\n            pred_df = infer_per_seed(cfg, model, data_loader, seed)\n            pred_dfs.append(pred_df)\n\n    return pl.concat(pred_dfs)\n\n\n@hydra.main(config_path=\"conf\", config_name=\"ensemble\", version_base=\"1.2\")\ndef main(cfg: EnsembleMainConfig):\n    data_dir = Path(cfg.env.data_dir)\n    working_dir = Path(cfg.env.working_dir)\n    eeg_dir = Path(working_dir / \"preprocess\" / cfg.phase / \"eeg\")\n    spec_dir = Path(working_dir / \"preprocess\" / cfg.phase / \"spectrogram\")\n    fold_split_dir = Path(working_dir / \"fold_split\" / cfg.phase)\n\n    metadata = load_metadata(\n        data_dir=data_dir,\n        phase=cfg.phase,\n        fold_split_dir=fold_split_dir,\n        num_samples=cfg.dev.num_samples,\n    )\n    with trace(\"** load eeg\"):\n        eeg_ids = metadata[\"eeg_id\"].unique().to_list()\n        id2eeg = preload_eegs(eeg_ids, eeg_dir)\n        id2cqf = preload_cqf(eeg_ids, eeg_dir)\n\n    with trace(\"** load bg spec\"):\n        spec_ids = metadata[\"spectrogram_id\"].unique().to_list()\n        spec_id2spec = preload_spectrograms(spec_ids, spec_dir)\n\n    with trace(\"** predict per experiments\"):\n        pred_dfs = []\n        for experiment in cfg.ensemble_entity.experiments:\n            print(f\"*** exp_name: {experiment.exp_name} ***\")\n            pred_df = infer_per_experiment(\n                cfg, experiment, metadata, id2eeg, id2cqf, spec_id2spec\n            )\n            pred_dfs.append(pred_df)\n            pr = (\n                pl.concat([pred_df])\n                .group_by(\"eeg_id\", maintain_order=True)\n                .agg(pl.col(f\"{label}_vote\").mean() for label in LABELS)\n            )\n\n\n            submission_df = make_submission(\n                eeg_ids=pr[\"eeg_id\"].to_list(),\n                predictions=pr.drop(\"eeg_id\").to_numpy(),\n                apply_softmax=False,\n            )\n            print(f'{len(submission_df)=}')\n            submission_dir = Path(cfg.env.submission_dir)\n            submission_df.write_csv(submission_dir / f\"{experiment.exp_name}.csv\")\n            print(submission_df)\n\n        pred_df = (\n            pl.concat(pred_dfs)\n            .group_by(\"eeg_id\", maintain_order=True)\n            .agg(pl.col(f\"{label}_vote\").mean() for label in LABELS)\n        )\n\n    with trace(\"** evaluate or make submission\"):\n        match cfg.phase:\n            case \"train\":\n                metadata = load_metadata(\n                    data_dir=data_dir,\n                    phase=cfg.phase,\n                    fold_split_dir=fold_split_dir,\n                    group_by_eeg=True,\n                    weight_key=\"weight_per_eeg\",\n                    num_samples=cfg.dev.num_samples,\n                )\n                do_evaluate(metadata, pred_df)\n\n            case \"test\":\n                submission_df = make_submission(\n                    eeg_ids=pred_df[\"eeg_id\"].to_list(),\n                    predictions=pred_df.drop(\"eeg_id\").to_numpy(),\n                    apply_softmax=False,\n                )\n                print(f'{len(submission_df)=}')\n                submission_dir = Path(cfg.env.submission_dir)\n                submission_df.write_csv(submission_dir / \"submission_bilzard.csv\")\n                print(submission_df)\n\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:52:34.515313Z","iopub.execute_input":"2024-04-14T05:52:34.515602Z","iopub.status.idle":"2024-04-14T05:52:34.530342Z","shell.execute_reply.started":"2024-04-14T05:52:34.515579Z","shell.execute_reply":"2024-04-14T05:52:34.529558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! python -m run.preprocess job_name=preprocess phase=test env=kaggle","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:52:34.531339Z","iopub.execute_input":"2024-04-14T05:52:34.531586Z","iopub.status.idle":"2024-04-14T05:52:47.406112Z","shell.execute_reply.started":"2024-04-14T05:52:34.531564Z","shell.execute_reply":"2024-04-14T05:52:47.405147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# if DEBUG:\n#     ! python -m run.preprocess job_name=preprocess phase=develop env=kaggle","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:52:47.407496Z","iopub.execute_input":"2024-04-14T05:52:47.407798Z","iopub.status.idle":"2024-04-14T05:52:47.412069Z","shell.execute_reply.started":"2024-04-14T05:52:47.407769Z","shell.execute_reply":"2024-04-14T05:52:47.411199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# if DEBUG:\n#     ! python -m run.batch_infer \\\n#         ensemble_entity={ENSEMBLE_ENTITY} \\\n#         ensemble_entity.name={ENSEMBLE_ENTITY_NAME} \\\n#         phase=develop \\\n#         env=kaggle","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:52:47.413244Z","iopub.execute_input":"2024-04-14T05:52:47.413529Z","iopub.status.idle":"2024-04-14T05:52:47.425132Z","shell.execute_reply.started":"2024-04-14T05:52:47.413507Z","shell.execute_reply":"2024-04-14T05:52:47.424298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! python -m run.new_batch_infer \\\n    ensemble_entity={ENSEMBLE_ENTITY} \\\n    ensemble_entity.name={ENSEMBLE_ENTITY_NAME} \\\n    phase=test \\\n    env=kaggle","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:52:47.426152Z","iopub.execute_input":"2024-04-14T05:52:47.426433Z","iopub.status.idle":"2024-04-14T05:54:24.824942Z","shell.execute_reply.started":"2024-04-14T05:52:47.426411Z","shell.execute_reply":"2024-04-14T05:54:24.823858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !rm /kaggle/working/submission_bilzard.csv\n!ls /kaggle/working/","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:24.826283Z","iopub.execute_input":"2024-04-14T05:54:24.826582Z","iopub.status.idle":"2024-04-14T05:54:25.762550Z","shell.execute_reply.started":"2024-04-14T05:54:24.826555Z","shell.execute_reply":"2024-04-14T05:54:25.761650Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%cd /kaggle/working","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:25.764055Z","iopub.execute_input":"2024-04-14T05:54:25.764446Z","iopub.status.idle":"2024-04-14T05:54:25.770956Z","shell.execute_reply.started":"2024-04-14T05:54:25.764409Z","shell.execute_reply":"2024-04-14T05:54:25.770180Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! rm -rf /kaggle/temp","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:25.772215Z","iopub.execute_input":"2024-04-14T05:54:25.772618Z","iopub.status.idle":"2024-04-14T05:54:26.728249Z","shell.execute_reply.started":"2024-04-14T05:54:25.772585Z","shell.execute_reply":"2024-04-14T05:54:26.727149Z"},"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":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# yu4u","metadata":{}},{"cell_type":"code","source":"use_yu4u_new_fold = False\n! cp /kaggle/input/hsm-yu4u-models/09_predict.py .","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:26.729781Z","iopub.execute_input":"2024-04-14T05:54:26.730150Z","iopub.status.idle":"2024-04-14T05:54:27.698854Z","shell.execute_reply.started":"2024-04-14T05:54:26.730112Z","shell.execute_reply":"2024-04-14T05:54:27.697738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! python /kaggle/input/hms-yu4u-old-split/09_predict.py --data_root /kaggle/input/hms-harmful-brain-activity-classification --mode test --checkpoint /kaggle/input/hms-yu4u-old-split --output_filename yu4u.csv\n# ! python /kaggle/input/hms-yu4u-old-split-stem/09_predict.py --data_root /kaggle/input/hms-harmful-brain-activity-classification --mode test --checkpoint /kaggle/input/hms-yu4u-old-split-stem --output_filename 1d_old_split_wd1e-2_stem_bs192_syncbn_dp0.2_ep128.csv","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:27.700587Z","iopub.execute_input":"2024-04-14T05:54:27.701196Z","iopub.status.idle":"2024-04-14T05:54:36.429082Z","shell.execute_reply.started":"2024-04-14T05:54:27.701148Z","shell.execute_reply":"2024-04-14T05:54:36.427861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ariyasu","metadata":{}},{"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\n        \nseed_everything()\ndebug = len(pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/test.csv')) == 1\ndebug = False\nDEVICE = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:36.436083Z","iopub.execute_input":"2024-04-14T05:54:36.436403Z","iopub.status.idle":"2024-04-14T05:54:38.575159Z","shell.execute_reply.started":"2024-04-14T05:54:36.436374Z","shell.execute_reply":"2024-04-14T05:54:38.573827Z"},"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-14T05:54:38.576615Z","iopub.execute_input":"2024-04-14T05:54:38.577070Z","iopub.status.idle":"2024-04-14T05:54:38.608227Z","shell.execute_reply.started":"2024-04-14T05:54:38.577028Z","shell.execute_reply":"2024-04-14T05:54:38.607110Z"},"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","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:38.611966Z","iopub.execute_input":"2024-04-14T05:54:38.612308Z","iopub.status.idle":"2024-04-14T05:54:41.991357Z","shell.execute_reply.started":"2024-04-14T05:54:38.612277Z","shell.execute_reply":"2024-04-14T05:54:41.990294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict_one_fmax(config_list, df):\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            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}.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-14T05:54:41.993036Z","iopub.execute_input":"2024-04-14T05:54:41.993368Z","iopub.status.idle":"2024-04-14T05:54:42.011988Z","shell.execute_reply.started":"2024-04-14T05:54:41.993340Z","shell.execute_reply":"2024-04-14T05:54:42.011053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# dataset","metadata":{}},{"cell_type":"code","source":"def 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-14T05:54:42.013608Z","iopub.execute_input":"2024-04-14T05:54:42.013875Z","iopub.status.idle":"2024-04-14T05:54:42.033799Z","shell.execute_reply.started":"2024-04-14T05:54:42.013854Z","shell.execute_reply":"2024-04-14T05:54:42.033062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if debug:\n    df = pd.read_csv('/kaggle/input/hms-data/train_unique_egg_v2.csv').query('fold==0')#[['spectrogram_id', 'eeg_id', 'patient_id']]\n    \n    edf = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/train.csv').drop_duplicates(['eeg_id','seizure_vote','lpd_vote','gpd_vote','lrda_vote','grda_vote','other_vote'])\n    edf['n'] = edf[['seizure_vote','lpd_vote','gpd_vote','lrda_vote','grda_vote','other_vote']].sum(1)\n#     ids = edf[edf['n']>8].eeg_id    \n    ids = edf[edf['n'].isin([10,11,12,13])].eeg_id    \n    \n    df = df[df.eeg_id.isin(ids)]\n    \n    train_test = 'train'\nelse:\n    df = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/test.csv')\n    train_test = 'test'\ndf['eeg_path'] = f'/kaggle/input/hms-harmful-brain-activity-classification/{train_test}_eeg/'+df.eeg_id.astype(str)+'.parquet'\ndf['spe_path'] = f'/kaggle/input/hms-harmful-brain-activity-classification/{train_test}_spectrograms/'+df.spectrogram_id.astype(str)+'.parquet'\nprint(len(df))","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:42.034886Z","iopub.execute_input":"2024-04-14T05:54:42.035173Z","iopub.status.idle":"2024-04-14T05:54:42.053892Z","shell.execute_reply.started":"2024-04-14T05:54:42.035141Z","shell.execute_reply":"2024-04-14T05:54:42.053054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# kaggle spectrograms","metadata":{}},{"cell_type":"code","source":"import albumentations as A\nclass hms_base_kaggle_spe():\n    def __init__(self):\n        super().__init__()\n        self.model_name = 'convnext_small.fb_in22k_ft_in1k_384'\n        self.image_size = (304, 400)\n        self.label_features = ['seizure','lpd','gpd','lrda','grda','other']\n        self.num_classes = len(self.label_features)\n        self.model = timm.create_model(self.model_name, pretrained=False, num_classes=self.num_classes)\n        self.transform = A.Compose([A.Resize(self.image_size[0], self.image_size[1])])\n        self.batch_size = 4\n        self.use_center_sec = None\n        self.tta = None\n        self.kaggle_spe = True\n\nclass fixfinetune_hms_swin_v3_mixup2(hms_base_kaggle_spe):\n    def __init__(self):\n        super().__init__()\n        self.model_name = 'swinv2_base_window12to24_192to384.ms_in22k_ft_in1k'\n        self.model = timm.create_model(self.model_name, pretrained=False, num_classes=self.num_classes)\n        if debug:\n            self.model_paths = [f'/kaggle/input/d/yujiariyasu/fixfinetune-hms-swin-v3-mixup2/last_fold{fold}.ckpt' for fold in [0]]\n        else:\n            self.model_paths = [f'/kaggle/input/d/yujiariyasu/fixfinetune-hms-swin-v3-mixup2/last_fold{fold}.ckpt' for fold in range(5)]\n        self.image_size = (384, 384)\n        self.transform = A.Compose([A.Resize(self.image_size[0], self.image_size[1])])\n        self.tta = True\nclass hms_swin_origin_label_20k(hms_base_kaggle_spe):\n    def __init__(self):\n        super().__init__()\n        self.model_name = 'swinv2_base_window12to24_192to384.ms_in22k_ft_in1k'\n        self.model = timm.create_model(self.model_name, pretrained=False, num_classes=self.num_classes)\n        d = 'hms-swin-origin-label-20k'\n        if debug:\n            self.model_paths = [f'/kaggle/input/d/yujiariyasu/{d}/last_fold{fold}.ckpt' for fold in [0]]\n        else:\n            self.model_paths = [f'/kaggle/input/d/yujiariyasu/{d}/last_fold{fold}.ckpt' for fold in range(5)]\n        self.image_size = (384, 384)\n        self.transform = A.Compose([A.Resize(self.image_size[0], self.image_size[1])])\n        self.tta = True\n\nclass fixfinetune_hms_swin_unique_spe_epoch10(hms_base_kaggle_spe):\n    def __init__(self):\n        super().__init__()\n        self.model_name = 'swinv2_base_window12to24_192to384.ms_in22k_ft_in1k'\n        self.model = timm.create_model(self.model_name, pretrained=False, num_classes=self.num_classes)\n        if debug:\n            self.model_paths = [f'/kaggle/input/d/yujiariyasu/fixfinetune-hms-swin-unique-spe-epoch10/last_fold{fold}.ckpt' for fold in [0]]\n        else:\n            self.model_paths = [f'/kaggle/input/d/yujiariyasu/fixfinetune-hms-swin-unique-spe-epoch10/last_fold{fold}.ckpt' for fold in range(5)]\n        self.image_size = (384, 384)\n        self.transform = A.Compose([A.Resize(self.image_size[0], self.image_size[1])])        \n        \nclass fixfinetune_hms_v23_same_mu(hms_base_kaggle_spe):\n    def __init__(self):\n        super().__init__()\n        self.model_name = 'convnext_small.fb_in22k_ft_in1k_384'\n        self.image_size = (304, 400)\n        self.model = timm.create_model(self.model_name, pretrained=False, num_classes=self.num_classes)\n        d = '/kaggle/input/d/yujiariyasu/fixfinetune-hms-v23-same-mu'\n        if debug:\n            self.model_paths = [f'{d}/last_fold{fold}.ckpt' for fold in [0]]\n        else:\n            self.model_paths = [f'{d}/last_fold{fold}.ckpt' for fold in range(5)]\n        self.transform = A.Compose([A.Resize(self.image_size[0], self.image_size[1])])        \n        \n            \nconfig_list = [\n    'fixfinetune_hms_swin_v3_mixup2',\n    'fixfinetune_hms_swin_unique_spe_epoch10',\n]\n","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:42.054924Z","iopub.execute_input":"2024-04-14T05:54:42.055229Z","iopub.status.idle":"2024-04-14T05:54:43.412899Z","shell.execute_reply.started":"2024-04-14T05:54:42.055205Z","shell.execute_reply":"2024-04-14T05:54:43.411921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# predict_one_fmax(config_list, df)","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:43.414003Z","iopub.execute_input":"2024-04-14T05:54:43.414459Z","iopub.status.idle":"2024-04-14T05:54:43.422013Z","shell.execute_reply.started":"2024-04-14T05:54:43.414432Z","shell.execute_reply":"2024-04-14T05:54:43.421257Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 549 fixfinetune_hms_swin_v3_mixup2 0.3303","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:43.423310Z","iopub.execute_input":"2024-04-14T05:54:43.423679Z","iopub.status.idle":"2024-04-14T05:54:43.434095Z","shell.execute_reply.started":"2024-04-14T05:54:43.423648Z","shell.execute_reply":"2024-04-14T05:54:43.433269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# egg spectrogram","metadata":{}},{"cell_type":"code","source":"PATH = f'/kaggle/input/hms-harmful-brain-activity-classification/{train_test}_eegs/'\nEEG_IDS = df.eeg_id.unique()\nUSE_WAVELET = None\nDISPLAY = 0","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:43.435151Z","iopub.execute_input":"2024-04-14T05:54:43.435422Z","iopub.status.idle":"2024-04-14T05:54:43.443801Z","shell.execute_reply.started":"2024-04-14T05:54:43.435392Z","shell.execute_reply":"2024-04-14T05:54:43.443022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"directory_path = f'/kaggle/temp/EEG_Spectrograms'\nif not os.path.exists(directory_path):\n    os.makedirs(directory_path)","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:43.444729Z","iopub.execute_input":"2024-04-14T05:54:43.444984Z","iopub.status.idle":"2024-04-14T05:54:43.453432Z","shell.execute_reply.started":"2024-04-14T05:54:43.444962Z","shell.execute_reply":"2024-04-14T05:54:43.452701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['path'] = '/kaggle/temp/EEG_Spectrograms/' + df.eeg_id.astype(str) + '.npy'","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:43.454432Z","iopub.execute_input":"2024-04-14T05:54:43.454713Z","iopub.status.idle":"2024-04-14T05:54:43.464284Z","shell.execute_reply.started":"2024-04-14T05:54:43.454681Z","shell.execute_reply":"2024-04-14T05:54:43.463488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import librosa\nfrom tqdm import tqdm\nimport pywt\nimport multiprocessing\nfrom scipy.signal import butter, lfilter\n\ndef maddest(d, axis=None):\n    return np.mean(np.absolute(d - np.mean(d, axis)), axis)\n\ndef denoise(x, wavelet='haar', level=1):\n    coeff = pywt.wavedec(x, wavelet, mode=\"per\")\n    sigma = (1/0.6745) * maddest(coeff[-level])\n\n    uthresh = sigma * np.sqrt(2*np.log(len(x)))\n    coeff[1:] = (pywt.threshold(i, value=uthresh, mode='hard') for i in coeff[1:])\n\n    ret=pywt.waverec(coeff, wavelet, mode='per')\n\n    return ret\n\n\ndef butter_bandpass(lowcut, highcut, fs, order=2):\n    nyq = 0.5 * fs\n    low = lowcut / nyq\n    high = highcut / nyq\n    b, a = butter(order, [low, high], btype=\"band\")\n    return b, a\n\n\ndef butter_bandpass_filter(data, lowcut=0.5, highcut=30, fs=200, order=2):\n    b, a = butter_bandpass(lowcut, highcut, fs, order=order)\n    y = lfilter(b, a, data)\n    return y","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:43.465208Z","iopub.execute_input":"2024-04-14T05:54:43.465513Z","iopub.status.idle":"2024-04-14T05:54:43.480551Z","shell.execute_reply.started":"2024-04-14T05:54:43.465489Z","shell.execute_reply":"2024-04-14T05:54:43.479635Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import librosa\n!mkdir -p /kaggle/temp\ndef spectrogram_from_eeg(args):\n    eeg_id, parquet_path, fmax, win_length, hop, im_height, w_cut, display, sec, apply_butter_bandpass_filter = args\n    eeg = pd.read_parquet(parquet_path)\n    t = sec*200\n    middle = (len(eeg)-t)//2\n    eeg = eeg.iloc[middle:middle+t]\n    img = np.zeros((128,im_height,len(FEATS)),dtype='float32')\n    \n#     if w_cut:\n#         img = np.zeros((128,im_height,len(FEATS)),dtype='float32')\n#     else:\n#         if hop == 512:\n#             img = np.zeros((128,im_height,len(FEATS)),dtype='float32')\n#         else:\n#             raise\n    \n    for k in range(len(FEATS)):\n        COLS = FEATS[k]\n        \n        for kk in range(len(COLS)-1):\n        \n            # COMPUTE PAIR DIFFERENCES\n            x = eeg[COLS[kk]].values - eeg[COLS[kk+1]].values\n\n            # FILL NANS\n            m = np.nanmean(x)\n            if np.isnan(x).mean()<1: x = np.nan_to_num(x,nan=m)\n            else: x[:] = 0\n\n            if USE_WAVELET:\n                x = denoise(x, wavelet=USE_WAVELET)\n            if apply_butter_bandpass_filter:\n                x = butter_bandpass_filter(x, lowcut=0.5, highcut=fmax, fs=200, order=2)\n            if display:\n                print(f'{fmax=}, {win_length=}, {hop=}')\n\n            mel_spec = librosa.feature.melspectrogram(y=x, sr=200, hop_length=len(x)//hop, \n                  n_fft=1024, n_mels=128, fmin=0, fmax=fmax, win_length=win_length)\n\n            width = (mel_spec.shape[1]//32)*32\n            if w_cut:\n                mel_spec = mel_spec[:, :width]\n            if (k==0) & (kk==0):\n                img = np.zeros((mel_spec.shape[0], mel_spec.shape[1], len(FEATS)), dtype='float32')\n\n            img[:,:,k] += mel_spec\n                \n        img[:,:,k] /= (len(COLS)-1)\n        \n    img = np.log(img+1e-9)       \n    np.save(f'{directory_path}/{eeg_id}',img)\n        \n#     return img\n\n\ndef eeg_to_spectrograms(fmax, hop=256, win_length=128, im_height=256, sec=50, apply_butter_bandpass_filter=False):\n    w_cut = True\n    display = False\n    args = [(eeg_id, f'{PATH}{eeg_id}.parquet', fmax, win_length, hop, im_height, w_cut, display, sec, apply_butter_bandpass_filter) for eeg_id in EEG_IDS]\n    print('eeg to spectrograms...')\n    p = multiprocessing.Pool(4)\n    p.map(spectrogram_from_eeg, args)\n    p.close()\n    if debug:\n        im = np.load('/kaggle/temp/EEG_Spectrograms/10617205.npy')\n        print(f'10617205.npy: {im.shape=}, {im.mean()=}') # 7.4906907","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:43.481938Z","iopub.execute_input":"2024-04-14T05:54:43.482225Z","iopub.status.idle":"2024-04-14T05:54:44.450963Z","shell.execute_reply.started":"2024-04-14T05:54:43.482201Z","shell.execute_reply":"2024-04-14T05:54:44.449712Z"},"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-14T05:54:44.453038Z","iopub.execute_input":"2024-04-14T05:54:44.453961Z","iopub.status.idle":"2024-04-14T05:54:44.460821Z","shell.execute_reply.started":"2024-04-14T05:54:44.453917Z","shell.execute_reply":"2024-04-14T05:54:44.459939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 4ims","metadata":{}},{"cell_type":"code","source":"FEATS = [['Fp1','F7','T3','T5','O1'],\n         ['Fp1','F3','C3','P3','O1'],\n         ['Fp2','F8','T4','T6','O2'],\n         ['Fp2','F4','C4','P4','O2']]","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:44.462202Z","iopub.execute_input":"2024-04-14T05:54:44.462650Z","iopub.status.idle":"2024-04-14T05:54:44.474465Z","shell.execute_reply.started":"2024-04-14T05:54:44.462617Z","shell.execute_reply":"2024-04-14T05:54:44.473579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # fmax60, win_length256\n# class finetune_hms_chris_fmax60_win256_v2_swinv2_large_spe_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#         if debug:\n#             self.model_paths = [f'/kaggle/input/finetune-fmax60-win256-swinv2-large-spe-eeg/last_fold{fold}.ckpt' for fold in [0]]\n#         else:\n#             self.model_paths = [f'/kaggle/input/finetune-fmax60-win256-swinv2-large-spe-eeg/last_fold{fold}.ckpt' for fold in range(5)]\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.4706545, 3.4938297)\n#         self.spe_eeg_concat = True\n#         self.skip_add_im_resize = True\n\n# config_list = [\n#     'finetune_hms_chris_fmax60_win256_v2_swinv2_large_spe_eeg'\n# ]\n# eeg_to_spectrograms(fmax=60, hop=512, win_length=256, sec=50)\n# predict_one_fmax(config_list, df)","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:44.475581Z","iopub.execute_input":"2024-04-14T05:54:44.475828Z","iopub.status.idle":"2024-04-14T05:54:44.485200Z","shell.execute_reply.started":"2024-04-14T05:54:44.475806Z","shell.execute_reply":"2024-04-14T05:54:44.484388Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 8ims","metadata":{}},{"cell_type":"code","source":"FEATS = [\n    ['Fp1','F7','T3'],\n    ['T3','T5','O1'],\n    ['Fp1','F3','C3'],\n    ['C3','P3','O1'],\n    ['Fp2','F8','T4'],\n    ['T4','T6','O2'],\n    ['Fp2','F4','C4'],\n    ['C4','P4','O2'],\n]","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:44.486125Z","iopub.execute_input":"2024-04-14T05:54:44.486397Z","iopub.status.idle":"2024-04-14T05:54:44.494935Z","shell.execute_reply.started":"2024-04-14T05:54:44.486374Z","shell.execute_reply":"2024-04-14T05:54:44.494141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# fmax30, 30sec, bandpass\n\nclass newdata_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg_9ep(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        if debug:\n            self.model_paths = [\n                '/kaggle/input/fold012-fmax30-30sec-8ims-2stage-pretrain/hms/results/finetune_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg/last_fold0.ckpt',\n            ]\n        else:\n            self.model_paths = [\n                '/kaggle/input/fold012-fmax30-30sec-8ims-2stage-pretrain/hms/results/finetune_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg/last_fold0.ckpt',\n                '/kaggle/input/fold012-fmax30-30sec-8ims-2stage-pretrain/hms/results/finetune_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg/last_fold1.ckpt',\n                '/kaggle/input/fold012-fmax30-30sec-8ims-2stage-pretrain/hms/results/finetune_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg/last_fold2.ckpt',\n                '/kaggle/input/fold34-fmax30-30sec-8ims-2stage-pretrain/hms/results/finetune_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg/last_fold3.ckpt',\n                '/kaggle/input/fold34-fmax30-30sec-8ims-2stage-pretrain/hms/results/finetune_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg/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        self.spe_eeg_concat = True\n\nconfig_list = [\n#     'newdata_hms_chris_fmax30_30sec_8ims_bandpass',\n    'newdata_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg_9ep'\n]\neeg_to_spectrograms(fmax=30, hop=256, win_length=128, im_height=544, sec=30, apply_butter_bandpass_filter=True)\npredict_one_fmax(config_list, df)\n","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:54:44.496013Z","iopub.execute_input":"2024-04-14T05:54:44.496294Z","iopub.status.idle":"2024-04-14T05:55:47.108143Z","shell.execute_reply.started":"2024-04-14T05:54:44.496265Z","shell.execute_reply":"2024-04-14T05:55:47.107108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# fmax90, 10sec, bandpass\nclass newdata_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg_9ep(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        c = 'hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg'\n        d = f\"/kaggle/input/{c.replace('_','-')}\"\n        if debug:\n            self.model_paths = [\n            '/kaggle/input/fold012-fmax90-10sec-2stage-pretrain/hms/results/finetune_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg/last_fold0.ckpt',\n            ]\n        else:\n            self.model_paths = [\n            '/kaggle/input/fold012-fmax90-10sec-2stage-pretrain/hms/results/finetune_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg/last_fold0.ckpt',\n            '/kaggle/input/fold012-fmax90-10sec-2stage-pretrain/hms/results/finetune_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg/last_fold1.ckpt',\n            '/kaggle/input/fold012-fmax90-10sec-2stage-pretrain/hms/results/finetune_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg/last_fold2.ckpt',\n            '/kaggle/input/fold34-fmax90-10sec-2stage-pretrain/hms/results/finetune_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg/last_fold3.ckpt',\n            '/kaggle/input/fold34-fmax90-10sec-2stage-pretrain/hms/results/finetune_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg/last_fold4.ckpt',\n\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        self.spe_eeg_concat = True\n\nconfig_list = [\n    'newdata_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg_9ep'\n]\n\neeg_to_spectrograms(fmax=90, hop=256, win_length=128, im_height=640, sec=10, apply_butter_bandpass_filter=True)\npredict_one_fmax(config_list, df)\n","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:55:47.110502Z","iopub.execute_input":"2024-04-14T05:55:47.110822Z","iopub.status.idle":"2024-04-14T05:56:43.348706Z","shell.execute_reply.started":"2024-04-14T05:55:47.110791Z","shell.execute_reply":"2024-04-14T05:56:43.347652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 16ims","metadata":{}},{"cell_type":"code","source":"FEATS = [\n['Fp1','F7'],\n['F7','T3'],\n['T3','T5'],\n['T5','O1'],\n['Fp1','F3'],\n['F3','C3'],\n['C3','P3'],\n['P3','O1'],\n['Fp2','F8'],\n['F8','T4'],\n['T4','T6'],\n['T6','O2'],\n['Fp2','F4'],\n['F4','C4'],\n['C4','P4'],\n['P4','O2'],\n]","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:56:43.351199Z","iopub.execute_input":"2024-04-14T05:56:43.352188Z","iopub.status.idle":"2024-04-14T05:56:43.358310Z","shell.execute_reply.started":"2024-04-14T05:56:43.352150Z","shell.execute_reply":"2024-04-14T05:56:43.357200Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\nclass newdata_hms_chris_fmax60_40sec_16ims_bandpass_9ep(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        if debug:\n            self.model_paths = [\n                '/kaggle/input/fold0-fmax60-40sec-16ims-12epochs/hms/results/finetune_hms_chris_fmax60_40sec_16ims_bandpass/last_fold0.ckpt',\n            ]\n        else:\n            self.model_paths = [\n                '/kaggle/input/fold0-fmax60-40sec-16ims-12epochs/hms/results/finetune_hms_chris_fmax60_40sec_16ims_bandpass/last_fold0.ckpt',\n                '/kaggle/input/fold1-fmax60-40sec-16ims-12epochs/hms/results/finetune_hms_chris_fmax60_40sec_16ims_bandpass/last_fold1.ckpt',\n                '/kaggle/input/fold2-fmax60-40sec-16ims-12epochs/hms/results/finetune_hms_chris_fmax60_40sec_16ims_bandpass/last_fold2.ckpt',\n                '/kaggle/input/fold3-fmax60-40sec-16ims-12epochs-e60a77/hms/results/finetune_hms_chris_fmax60_40sec_16ims_bandpass/last_fold3.ckpt',\n                '/kaggle/input/fold4-fmax60-40sec-16ims-12epochs-a8660e/hms/results/finetune_hms_chris_fmax60_40sec_16ims_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\nconfig_list = [\n    'newdata_hms_chris_fmax60_40sec_16ims_bandpass_9ep'\n]\neeg_to_spectrograms(fmax=60, hop=256, win_length=128, im_height=256, sec=40, apply_butter_bandpass_filter=True)\npredict_one_fmax(config_list, df)","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:56:43.359674Z","iopub.execute_input":"2024-04-14T05:56:43.360262Z","iopub.status.idle":"2024-04-14T05:57:41.410160Z","shell.execute_reply.started":"2024-04-14T05:56:43.360213Z","shell.execute_reply":"2024-04-14T05:57:41.409038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# fmax30, 50sec, bandpass\nclass newdata_hms_chris_fmax30_50sec_16ims_bandpass_9ep(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        if debug:\n            self.model_paths = [\n                '/kaggle/input/fold0-fmax30-50sec-9epochs-2stage-pretrain/hms/results/finetune_hms_chris_fmax30_50sec_16ims_bandpass/last_fold0.ckpt',\n            ]\n        else:\n            self.model_paths = [\n                '/kaggle/input/fold0-fmax30-50sec-9epochs-2stage-pretrain/hms/results/finetune_hms_chris_fmax30_50sec_16ims_bandpass/last_fold0.ckpt',\n                '/kaggle/input/fold12-fmax30-50sec-9epochs-2stage-pretrain/hms/results/finetune_hms_chris_fmax30_50sec_16ims_bandpass/last_fold1.ckpt',\n                '/kaggle/input/fold12-fmax30-50sec-9epochs-2stage-pretrain/hms/results/finetune_hms_chris_fmax30_50sec_16ims_bandpass/last_fold2.ckpt',\n                '/kaggle/input/fold34-fmax30-50sec-9epochs-2stage-pretrai/hms/results/finetune_hms_chris_fmax30_50sec_16ims_bandpass/last_fold3.ckpt',\n                '/kaggle/input/fold34-fmax30-50sec-9epochs-2stage-pretrai/hms/results/finetune_hms_chris_fmax30_50sec_16ims_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\nconfig_list = [\n    'newdata_hms_chris_fmax30_50sec_16ims_bandpass_9ep'\n]\neeg_to_spectrograms(fmax=30, hop=256, win_length=128, im_height=256, sec=50, apply_butter_bandpass_filter=True)\npredict_one_fmax(config_list, df)","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:57:41.412313Z","iopub.execute_input":"2024-04-14T05:57:41.412627Z","iopub.status.idle":"2024-04-14T05:58:39.233971Z","shell.execute_reply.started":"2024-04-14T05:57:41.412596Z","shell.execute_reply":"2024-04-14T05:58:39.232805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if debug:\n    configs = [\n        'newdata_hms_chris_fmax30_50sec_16ims_bandpass_9ep',\n        'newdata_hms_chris_fmax60_40sec_16ims_bandpass_9ep',\n        'newdata_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg_9ep',\n        'newdata_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg_9ep',\n    ]\n    \n    preds = []\n    true_cols = ['seizure','lpd','gpd','lrda','grda','other']\n    weights = [0.32785061, 0.20382992, 0.24517077, 0.22314871]\n    assert len(configs) == len(weights)\n    for w, c in zip(weights, configs):\n        pred_df = pd.read_csv(f'pred_{c}.csv').sort_values('eeg_id')\n        pred_cols = [f'{c}_pred_{col}' for col in true_cols]\n        preds.append(pred_df[pred_cols].values*w)\n    \n    oof = df.copy().sort_values('eeg_id')\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.sum(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-14T05:58:39.236051Z","iopub.execute_input":"2024-04-14T05:58:39.236393Z","iopub.status.idle":"2024-04-14T05:58:39.246261Z","shell.execute_reply.started":"2024-04-14T05:58:39.236362Z","shell.execute_reply":"2024-04-14T05:58:39.245400Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -rf /kaggle/temp","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:58:39.247531Z","iopub.execute_input":"2024-04-14T05:58:39.247867Z","iopub.status.idle":"2024-04-14T05:58:40.279671Z","shell.execute_reply.started":"2024-04-14T05:58:39.247835Z","shell.execute_reply":"2024-04-14T05:58:40.278348Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"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":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# use_yu4u_new_fold=False","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:58:40.281371Z","iopub.execute_input":"2024-04-14T05:58:40.281711Z","iopub.status.idle":"2024-04-14T05:58:40.286173Z","shell.execute_reply.started":"2024-04-14T05:58:40.281674Z","shell.execute_reply":"2024-04-14T05:58:40.285296Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# mlp","metadata":{}},{"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-14T05:58:40.287742Z","iopub.execute_input":"2024-04-14T05:58:40.288325Z","iopub.status.idle":"2024-04-14T05:58:40.304082Z","shell.execute_reply.started":"2024-04-14T05:58:40.288292Z","shell.execute_reply":"2024-04-14T05:58:40.303277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%cd /kaggle/working","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:58:40.305398Z","iopub.execute_input":"2024-04-14T05:58:40.305978Z","iopub.status.idle":"2024-04-14T05:58:40.319031Z","shell.execute_reply.started":"2024-04-14T05:58:40.305947Z","shell.execute_reply":"2024-04-14T05:58:40.318165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 一旦 5939件\ncols = ['seizure_vote','lpd_vote','gpd_vote','lrda_vote','grda_vote','other_vote']\ndf = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/train.csv')\ndf['n'] = df[cols].sum(1)\ndf = df[df.n>8]\ndfs = []\nfor i, idf in df.groupby('eeg_id'):\n    idf[cols] = idf[cols].mean(0)\n    dfs.append(idf.iloc[:1])\ndf = pd.concat(dfs)\ndf[cols] = df[cols] / np.array([df[cols].values.sum(1).tolist()]*6).T\ntrue = df.sort_values('eeg_id')[cols].values+1e-10\nlen(true)","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:58:40.320271Z","iopub.execute_input":"2024-04-14T05:58:40.320571Z","iopub.status.idle":"2024-04-14T05:59:03.203407Z","shell.execute_reply.started":"2024-04-14T05:58:40.320540Z","shell.execute_reply":"2024-04-14T05:59:03.202525Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ariyasu\nprint('ariyasu:')\nconfigs=[\n    'newdata_hms_chris_fmax30_50sec_16ims_bandpass_9ep',\n    'newdata_hms_chris_fmax60_40sec_16ims_bandpass_9ep', # pretrain twice\n    'newdata_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg_9ep',\n    'newdata_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg_9ep',\n\n#     'newdata_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg',\n#     'newdata_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg',\n        \n]\ntrue_cols = ['seizure','lpd','gpd','lrda','grda','other']\npred_cols = [f'pred_{c}' for c in true_cols]\npreds = []\nfolds = range(5)\nfor config in configs:\n    if os.path.exists(f'/kaggle/input/oof-by-kaggle-note/{config}.csv'):\n        oof = pd.read_csv(f'/kaggle/input/oof-by-kaggle-note/{config}.csv').sort_values('eeg_id')\n        oof = oof[oof.eeg_id.isin(clean_eeg_ids)].sort_values('eeg_id')\n        expert_consensus = oof['expert_consensus'].values\n        pr = oof[pred_cols].values\n    else:\n        if config == 'newdata_hms_chris_fmax30_50sec_16ims_bandpass_9ep':\n#             oof = pd.read_csv('/kaggle/input/ariyasu-newdata-oof-v2/finetune_hms_chris_fmax30_50sec_16ims_bandpass.csv')\n#             cs = [f'finetune_hms_chris_fmax30_50sec_16ims_bandpass_pred_{c}' for c in true_cols]\n            oof = pd.read_csv('/kaggle/input/ariyasu-newdata-oof-v3/finetune_hms_chris_fmax30_50sec_16ims_bandpass_ep9.csv')\n            cs = [f'finetune_hms_chris_fmax30_50sec_16ims_bandpass_ep9_pred_{c}' for c in true_cols]\n        elif config == 'newdata_hms_chris_fmax60_40sec_16ims_bandpass_9ep':\n#             oof = pd.read_csv('/kaggle/input/ariyasu-newdata-oof-v2/finetune_hms_chris_fmax60_40sec_16ims_bandpass.csv')\n#             cs = [f'finetune_hms_chris_fmax60_40sec_16ims_bandpass_pred_{c}' for c in true_cols]\n            oof = pd.read_csv('/kaggle/input/fork-of-ariyasu-newdata-oof-v4/finetune_hms_chris_fmax60_40sec_16ims_bandpass_ep9.csv')\n            cs = [f'finetune_hms_chris_fmax60_40sec_16ims_bandpass_ep9_pred_{c}' for c in true_cols]\n        elif config == 'newdata_hms_chris_fmax30_50sec_16ims_bandpass':\n            oof = pd.read_csv('/kaggle/input/fork-of-finetune-hms-chris-fmax30-30sec-8ims-bandp/finetune_hms_chris_fmax30_50sec_16ims_bandpass.csv')\n            cs = [f'finetune_hms_chris_fmax30_50sec_16ims_bandpass_pred_{c}' for c in true_cols]\n        elif config == 'newdata_hms_chris_fmax30_30sec_8ims_bandpass':\n            oof = pd.read_csv('/kaggle/input/ariyasu-newdata-oof/newdata_finetune_hms_chris_fmax30_30sec_8ims_bandpass.csv')\n            cs = [f'finetune_hms_chris_fmax30_30sec_8ims_bandpass_pred_{c}' for c in true_cols]\n        elif config == 'newdata_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg':\n            oof = pd.read_csv('/kaggle/input/fork-of-finetune-hms-chris-fmax30-30sec-8ims-bandp/finetune_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg.csv')\n            cs = [f'finetune_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg_pred_{c}' for c in true_cols]\n        elif config == 'newdata_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg_9ep':\n            oof = pd.read_csv('/kaggle/input/ariyasu-newdata-oof-v4/finetune_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg_ep9.csv')\n            cs = [f'finetune_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg_ep9_pred_{c}' for c in true_cols]\n        elif config == 'newdata_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg_9ep':\n            oof = pd.read_csv('/kaggle/input/ariyasu-newdata-oof-v4/finetune_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg_ep9.csv')\n            cs = [f'finetune_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg_ep9_pred_{c}' for c in true_cols]\n        elif config == 'newdata_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg':\n            oof = pd.read_csv('/kaggle/input/fork-of-finetune-hms-chris-fmax30-30sec-8ims-bandp/finetune_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg.csv')\n            cs = [f'finetune_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg_pred_{c}' for c in true_cols]\n        oof = pd.DataFrame(oof.groupby('eeg_id').first()).reset_index().sort_values('eeg_id')\n        expert_consensus = oof['expert_consensus'].values\n        oof[true_cols] = true-1e-10\n        # oof = upsample(oof)\n        pr = oof[cs].values\n        oof[pred_cols] = pr\n    score = kl_divergence(oof[true_cols].values+1e-10, pr)\n    print(len(oof), config, round(score, 4))\n    preds.append(pr)\n\n# tattaka\nprint('tattaka:')\ncols = ['oof_logits_seizure_vote','oof_logits_lpd_vote','oof_logits_gpd_vote','oof_logits_lrda_vote','oof_logits_grda_vote','oof_logits_other_vote']\nexp_dirs = [\n#     \"/kaggle/input/hms-weights-2/exp092/tiny_vit_21m_384_el30_mixup_50ep\",\n#     \"/kaggle/input/hms-weights-2/exp094/caformer_s18_2_5d_256_el30_mixup_100ep\",\n#     \"/kaggle/input/hms-weights-2/exp108/convnext_large_384_el40_mixup_50ep\",    \n\n#     \"/kaggle/input/hms-weights-2/exp108/convnext_large_384_el40_mixup_50ep\",\n    \"/kaggle/input/hms-weights-2/exp094/caformer_s18_2_5d_256_el30_mixup_100ep\",\n    \"/kaggle/input/hms-weights-3/exp147/tiny_vit_21m_512_el30_mixup_50ep\",\n    \n]\nfor exp_dir in exp_dirs:\n    df = pd.read_csv(glob(f\"{exp_dir}/**/result_df.csv\", recursive=True)[0])\n    df = df[df.num_votes > 7].groupby(\"eeg_id\").first().reset_index()\n    c = exp_dir.split('/')[-1]\n    df = df[df.eeg_id.isin(oof.eeg_id)].sort_values('eeg_id')\n    pr = df[cols].values\n    preds.append(pr)\n    configs.append(c)\n    score = kl_divergence(true, pr)\n    print(c, round(score, 4))\nprint('-'*100)     \n    \n# bilzard\nprint('bilzard:')\ncols = ['pl_seizure_vote','pl_lpd_vote','pl_gpd_vote','pl_lrda_vote','pl_grda_vote','pl_other_vote']\nfor c in [\n# 'eeg022_16ep_sc03c',\n# 'v5_panns_12ep_sc03c',\n'v5_eeg_24ep_cutmix',\n# 'v5_spec_bg_8ep',\n]:\n    df = pd.read_parquet(f'/kaggle/input/hms-bilzard-oof/{c}/train_pseudo_label.pqt').sort_values('eeg_id')\n    df = df[df.eeg_id.isin(oof.eeg_id)].sort_values('eeg_id')\n    pr = df[cols].values\n    preds.append(pr)\n    configs.append(c)\n    score = kl_divergence(true, pr)\n    print(c, round(score, 4))\n    \nprint('-'*100)    \n\n# yu4u\nprint('yu4u:')\nfor c in [\n    \"yu4u\",\n#     \"1d_new_split_wd1e-2_bs128_syncbn_ds7\",\n#     \"1d_new_split_wd1e-2_stem_bs192_syncbn_dp0.2_ep128\",\n#     \"1d_old_split_wd1e-2_stem_bs192_syncbn_dp0.2_ep128\",\n#     \"1d_new_split_stem_ft2_1stage\",\n#     \"1d_new_split_wd1e-2_stem_bs128_syncbn_center\",\n]:\n    if c == \"yu4u\":\n        df = pd.read_csv(\"/kaggle/input/hms-oof/oof.csv\") # old fold\n#         c = \"yu4u_old\"\n    else:\n        df = pd.read_csv(f\"/kaggle/input/hms-oof/oof_result__{c}.csv\")\n    \n    pr = df.sort_values(\"eeg_id\")[[\"seizure_vote\", \"lpd_vote\", \"gpd_vote\", \"lrda_vote\", \"grda_vote\", \"other_vote\"]].values\n    preds.append(pr)\n    configs.append(c)\n    score = kl_divergence(true, pr)\n    print(c, round(score, 4))","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:59:03.204934Z","iopub.execute_input":"2024-04-14T05:59:03.205325Z","iopub.status.idle":"2024-04-14T05:59:04.406278Z","shell.execute_reply.started":"2024-04-14T05:59:03.205291Z","shell.execute_reply":"2024-04-14T05:59:04.405291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Nelder-Mead\nstarting_weights = [1/len(preds)] * len(preds)# * 6\nconstraints = ({'type': 'eq', 'fun': lambda w: 1-sum(w)})\nbounds = [(0, 1)] * (len(preds))\n# bounds = [(0, 1)] * (len(preds)*6)\nres = minimize(loss_fn, starting_weights, method='Nelder-Mead', bounds=bounds, constraints=constraints)\n","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:59:04.407285Z","iopub.execute_input":"2024-04-14T05:59:04.407550Z","iopub.status.idle":"2024-04-14T05:59:04.736432Z","shell.execute_reply.started":"2024-04-14T05:59:04.407527Z","shell.execute_reply":"2024-04-14T05:59:04.735666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"round(res['fun'], 4), res['x']","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:59:04.737473Z","iopub.execute_input":"2024-04-14T05:59:04.737758Z","iopub.status.idle":"2024-04-14T05:59:04.743699Z","shell.execute_reply.started":"2024-04-14T05:59:04.737734Z","shell.execute_reply":"2024-04-14T05:59:04.742862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# print\nweights = res['x'] / np.sum(res['x'])\nprint('weights:')\nfor w, c in zip(weights, configs):\n    print(c, round(w, 4))\n\nprint('best score:', round(res['fun'], 4))\n\nprint('='*100)\n\nrounded_weights = np.array([round_to_nearest(x) for x in weights])\nrounded_weights = rounded_weights / np.sum(rounded_weights)\nprint('rounded_weights:')\nfor w, c in zip(rounded_weights, configs):\n    print(c, round(w, 4))\n\nall_preds = []\nfor w, p in zip(rounded_weights, preds):\n    all_preds.append(p*w)\noof[pred_cols] = softmax(np.sum(all_preds, 0), 1)\n\ntrue_df = oof[true_cols].copy()\npred = oof[pred_cols].copy()\ntrue_df['id'] = list(range(len(true_df)))\npred['id'] = list(range(len(pred)))\ntrue_df.columns = [0,1,2,3,4,5,'id']\npred.columns = [0,1,2,3,4,5,'id']\nscore = kaggle_metric(solution=true_df, submission=pred, row_id_column_name='id')\nprint('score:', round(score, 4))","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:59:04.745075Z","iopub.execute_input":"2024-04-14T05:59:04.745391Z","iopub.status.idle":"2024-04-14T05:59:04.797938Z","shell.execute_reply.started":"2024-04-14T05:59:04.745367Z","shell.execute_reply":"2024-04-14T05:59:04.797108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"features = []\nassert len(configs) == len(preds)\nfor c, pr in zip(configs, preds):\n    oof[[f'{c}_{col}' for col in pred_cols]] = pr\n    features += [f'{c}_{col}' for col in pred_cols]\n\n# upsample\ndfs = []\nfor sample_n, v in zip(\n#     [4,3,3,4,4,3],\n    [11, 8, 8, 11, 11, 6],\n    ['Seizure','LPD','GPD','LRDA','GRDA','Other']\n):\n    for _ in range(sample_n):\n        dfs.append(oof[oof.expert_consensus==v].copy())\noof = pd.concat(dfs)\nlen(oof)","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:59:04.799304Z","iopub.execute_input":"2024-04-14T05:59:04.799544Z","iopub.status.idle":"2024-04-14T05:59:05.162909Z","shell.execute_reply.started":"2024-04-14T05:59:04.799522Z","shell.execute_reply":"2024-04-14T05:59:05.162023Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(preds))\np = f'oof_for_mlp.csv'\noof.to_csv(p, index=False)\nconfigs","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:59:05.164188Z","iopub.execute_input":"2024-04-14T05:59:05.164523Z","iopub.status.idle":"2024-04-14T05:59:10.767944Z","shell.execute_reply.started":"2024-04-14T05:59:05.164496Z","shell.execute_reply":"2024-04-14T05:59:10.767092Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"assert configs==['newdata_hms_chris_fmax30_50sec_16ims_bandpass_9ep',\n 'newdata_hms_chris_fmax60_40sec_16ims_bandpass_9ep',\n 'newdata_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg_9ep',\n 'newdata_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg_9ep',\n#  'convnext_large_384_el40_mixup_50ep',\n 'caformer_s18_2_5d_256_el30_mixup_100ep',\n 'tiny_vit_21m_512_el30_mixup_50ep',\n#  'eeg022_16ep_sc03c',\n 'v5_eeg_24ep_cutmix',\n 'yu4u',\n# \"1d_old_split_wd1e-2_stem_bs192_syncbn_dp0.2_ep128\"                \n                ]","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:59:10.768993Z","iopub.execute_input":"2024-04-14T05:59:10.769278Z","iopub.status.idle":"2024-04-14T05:59:10.774376Z","shell.execute_reply.started":"2024-04-14T05:59:10.769253Z","shell.execute_reply":"2024-04-14T05:59:10.773487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cp -r /kaggle/input/hms-pipeline ./","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:59:10.775926Z","iopub.execute_input":"2024-04-14T05:59:10.776264Z","iopub.status.idle":"2024-04-14T05:59:13.492046Z","shell.execute_reply.started":"2024-04-14T05:59:10.776216Z","shell.execute_reply":"2024-04-14T05:59:13.490725Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir -p /kaggle/working/hms/results","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:59:13.493879Z","iopub.execute_input":"2024-04-14T05:59:13.494294Z","iopub.status.idle":"2024-04-14T05:59:14.525930Z","shell.execute_reply.started":"2024-04-14T05:59:13.494255Z","shell.execute_reply":"2024-04-14T05:59:14.524643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%writefile /kaggle/working/hms-pipeline/hms_pipeline/src/mlp_configs.py","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:59:14.527390Z","iopub.execute_input":"2024-04-14T05:59:14.527679Z","iopub.status.idle":"2024-04-14T05:59:14.533426Z","shell.execute_reply.started":"2024-04-14T05:59:14.527652Z","shell.execute_reply":"2024-04-14T05:59:14.532550Z"},"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 in configs.py\")\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    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}-v*.ckpt')\n","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:59:14.536863Z","iopub.execute_input":"2024-04-14T05:59:14.537176Z","iopub.status.idle":"2024-04-14T05:59:14.555596Z","shell.execute_reply.started":"2024-04-14T05:59:14.537152Z","shell.execute_reply":"2024-04-14T05:59:14.554797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile /kaggle/working/hms-pipeline/hms_pipeline/src/lightning/lightning_modules/mlp.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\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        if self.cfg.mixup and (torch.rand(1)[0] < 0.5) and (self.cfg.warmup_epochs < self.current_epoch):\n            mix_images, target_a, target_b, lam = mixup(images, targets, alpha=0.5)\n            logits = self.forward(mix_images)\n            loss = self.cfg.criterion(logits, target_a) * lam + (1 - lam) * self.cfg.criterion(logits, target_b)\n        else:\n            logits = self.forward(images)\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        logits = self.forward(images)\n        if isinstance(logits, tuple):\n            logits = logits[0]\n        loss = self.cfg.criterion(logits, targets)\n        preds = logits\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    def on_validation_epoch_end(self):\n        outputs = self.validation_step_outputs\n\n        d = dict()\n        d[\"epoch\"] = int(self.current_epoch)\n\n        # targets = torch.cat([o[\"targets\"] for o in outputs]).float()\n        # preds = torch.cat([o[\"preds\"] for o in outputs]).float()\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\n        loss = self.cfg.criterion(preds, targets)\n        d[\"v_loss\"] = loss.item()\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.cpu(), preds.cpu())\n\n        d[\"val_metric\"] = score\n        print('val metric:', score)\n\n        self.log_dict(d, prog_bar=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","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:59:14.556817Z","iopub.execute_input":"2024-04-14T05:59:14.557079Z","iopub.status.idle":"2024-04-14T05:59:14.570796Z","shell.execute_reply.started":"2024-04-14T05:59:14.557056Z","shell.execute_reply":"2024-04-14T05:59:14.569908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile /kaggle/working/hms-pipeline/hms_pipeline/src/mlp_configs.py\n\nfrom pathlib import Path\nfrom pprint import pprint\nimport timm\nfrom src.utils.metrics import *\nfrom src.utils.loss import *\nimport os\nimport torch\nimport torch.nn as nn\nimport pandas as pd\nimport numpy as np\nfrom pdb import set_trace as st\n\nfrom sklearn.metrics import roc_auc_score, confusion_matrix, mean_squared_error, average_precision_score\nfrom src.models.mlp import *\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\n    def forward(self, logits, targets, weights=None):\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        return loss.mean()\ndef kl_divergence(true, preds):\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 Baseline:\n    def __init__(self):\n        self.compe = 'kaggle_days'\n        self.batch_size = 8\n        self.grad_accumulations = 1\n        self.lr = 0.0001\n        self.epochs = 300\n        self.resume = False\n        self.seed = 2022\n        self.tta = 1\n        self.predict_valid = True\n        self.predict_test = True\n        self.valid_df = None\n        self.num_classes = 1\n        self.criterion = torch.nn.BCEWithLogitsLoss()\n        # self.criterion = torch.nn.BCELoss()\n        self.label_features = ['target']\n        self.metric = roc_auc_score # AUC().torch # MultiAP().torch\n        self.fp16 = False\n        self.optimizer = 'adam'\n        # self.scheduler = 'linear_schedule_with_warmup'\n        self.scheduler = 'CosineAnnealingWarmRestarts'\n        self.train_by_all_data = False\n        self.early_stop_patience = 30\n        self.inference = False\n        self.logit_to = None\n        self.pretrained_path = None\n        self.sync_batchnorm = False\n        self.finetune_transform = None\n        self.gpu = 'v100'\n        self.weight_decay = 0.1 # | 0.01\n        self.inference_only = False\n        self.num_train_optimization_steps = 3000\n        self.resume_epoch = 0\n        self.t_max = 30\n        self.save_top_k = 1\n        self.output_features = False\n        self.eta_min = 5e-7\n        self.mixup = False\n        self.add_imsizes_when_inference = [(0, 0)] # dummy\n\n################\nclass hms_stacking_v4_ep500(Baseline):\n    def __init__(self):\n        super().__init__()\n        self.compe = 'hms'\n        self.train_df_path = '/kaggle/working/oof_for_mlp.csv'\n        self.train_df = pd.read_csv(self.train_df_path)\n        configs = [\n#         'newdata_hms_chris_fmax30_50sec_16ims_bandpass',    \n#         'newdata_hms_chris_fmax30_30sec_8ims_bandpass',\n#         'newdata_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg',\n#         'newdata_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg',            \n            \n    'newdata_hms_chris_fmax30_50sec_16ims_bandpass_9ep',\n    'newdata_hms_chris_fmax60_40sec_16ims_bandpass_9ep', # pretrain twice\n    'newdata_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg_9ep',\n    'newdata_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg_9ep',\n            \n#          'tiny_vit_21m_384_el30_mixup_50ep',\n#          'caformer_s18_2_5d_256_el30_mixup_100ep',\n#          'convnext_large_384_el40_mixup_50ep',\n#          'convnext_large_384_el40_mixup_50ep',\n         'caformer_s18_2_5d_256_el30_mixup_100ep',\n         'tiny_vit_21m_512_el30_mixup_50ep',            \n#          'eeg022_16ep_sc03c',\n         'v5_eeg_24ep_cutmix',\n         'yu4u',\n# \"1d_old_split_wd1e-2_stem_bs192_syncbn_dp0.2_ep128\"            \n        ]\n        print('len(configs):', len(configs))\n        self.label_features = ['seizure','lpd','gpd','lrda','grda','other']\n\n        self.meta_cols = []\n        for config in configs:\n            for col in self.label_features:\n                if f'{config}_pred_{col}' in list(self.train_df):\n                    self.meta_cols.append(f'{config}_pred_{col}')\n\n        self.num_classes = len(self.label_features)\n        self.model = MultiLayerPerceptronBase(num_classes=self.num_classes, input_num=len(self.meta_cols), dropout=0.3)\n        self.batch_size = 16\n        self.grad_accumulations = 2\n        self.lr = 1e-4\n        self.criterion = hms_criterion()\n        self.metric = kl_divergence\n        self.warmup_epochs = 0\n        self.predict_test = False\n        self.predict_valid = True\n        self.gpu = 'small'\n        self.valid_df = self.train_df.copy()\n        self.train_df['fold'] = -1\n        self.valid_df['fold'] = 0\n        self.t_max = 30\n        self.epochs = 500\n\nclass MultiLayerPerceptronBase2(nn.Module):\n    def __init__(self, num_classes, input_num):\n        super(MultiLayerPerceptronBase2, self).__init__()\n        config_n = input_num//6\n        self.mlps = []\n        for _ in range(config_n):\n            self.mlps.append(\n                nn.Linear(input_num//6, 1)\n            )\n\n    def forward(self, x, labels=None):\n        xs = []\n        for i, mlp in enumerate(self.mlps):\n            idxes = [j % 6 == i for j in range(x.size(1))]\n            idxes = torch.tensor(idxes).to(x.device)\n            class_x = x[:, idxes]\n            class_x = mlp(class_x)\n            xs.append(class_x)\n\n        xs = torch.cat(xs, dim=1)\n        return xs   \nclass MultiLayerPerceptronBase2(nn.Module):\n    def __init__(self, num_classes, input_num):\n        super(MultiLayerPerceptronBase2, self).__init__()\n        self.seizure = nn.Linear(input_num//6, 1)\n        self.lpd = nn.Linear(input_num//6, 1)\n        self.gpd = nn.Linear(input_num//6, 1)\n        self.lrda = nn.Linear(input_num//6, 1)\n        self.grda = nn.Linear(input_num//6, 1)\n        self.other = nn.Linear(input_num//6, 1)\n        self.mlps = [\n            self.seizure,\n            self.lpd,\n            self.gpd,\n            self.lrda,\n            self.grda,\n            self.other,\n        ]\n    def forward(self, x, labels=None):\n        xs = []\n        for i, mlp in enumerate(self.mlps):\n            idxes = [j for j in range(x.size(1)) if j % 6 == i]\n            idxes = torch.tensor(idxes).to(x.device)\n            class_x = x[:, idxes]\n            class_x = mlp(class_x)\n            xs.append(class_x)\n\n        xs = torch.cat(xs, dim=1)\n        return xs    \nclass MultiLayerPerceptronBaseNoBias2(nn.Module):\n    def __init__(self, num_classes, input_num):\n        super(MultiLayerPerceptronBaseNoBias2, self).__init__()\n        self.seizure = nn.Linear(input_num//6, 1, bias=False)\n        self.lpd = nn.Linear(input_num//6, 1, bias=False)\n        self.gpd = nn.Linear(input_num//6, 1, bias=False)\n        self.lrda = nn.Linear(input_num//6, 1, bias=False)\n        self.grda = nn.Linear(input_num//6, 1, bias=False)\n        self.other = nn.Linear(input_num//6, 1, bias=False)\n        self.mlps = [\n            self.seizure,\n            self.lpd,\n            self.gpd,\n            self.lrda,\n            self.grda,\n            self.other,\n        ]\n    def forward(self, x, labels=None):\n        xs = []\n        for i, mlp in enumerate(self.mlps):\n            idxes = [j % 6 == i for j in range(x.size(1))]\n            idxes = torch.tensor(idxes).to(x.device)\n            class_x = x[:, idxes]\n            class_x = mlp(class_x)\n            xs.append(class_x)\n\n        xs = torch.cat(xs, dim=1)\n        return xs    \n\nclass mlp_config_0(hms_stacking_v4_ep500):\n    def __init__(self):\n        super().__init__()\n        self.model = MultiLayerPerceptronBase(num_classes=self.num_classes, input_num=len(self.meta_cols))\n#         self.model = MultiLayerPerceptronBase2(num_classes=self.num_classes, input_num=len(self.meta_cols))\n        self.epochs = 500\nclass mlp_config_1(hms_stacking_v4_ep500):\n    def __init__(self):\n        super().__init__()\n        self.model = MultiLayerPerceptronBaseNoBias(num_classes=self.num_classes, input_num=len(self.meta_cols))\n#         self.model = MultiLayerPerceptronBaseNoBias2(num_classes=self.num_classes, input_num=len(self.meta_cols))\n        self.epochs = 500\n","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:59:14.572419Z","iopub.execute_input":"2024-04-14T05:59:14.572751Z","iopub.status.idle":"2024-04-14T05:59:14.588021Z","shell.execute_reply.started":"2024-04-14T05:59:14.572721Z","shell.execute_reply":"2024-04-14T05:59:14.587191Z"},"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-14T05:59:14.589027Z","iopub.execute_input":"2024-04-14T05:59:14.589308Z","iopub.status.idle":"2024-04-14T05:59:14.601702Z","shell.execute_reply.started":"2024-04-14T05:59:14.589280Z","shell.execute_reply":"2024-04-14T05:59:14.600874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nif len(pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/test.csv')) == 1:\n    !python3 train_one_fold.py -c mlp_config_0 -t mlp -e 1\n    !python3 train_one_fold.py -c mlp_config_1 -t mlp -e 1\nelse:\n    !python3 train_one_fold.py -c mlp_config_0 -t mlp -e 300\n    !python3 train_one_fold.py -c mlp_config_1 -t mlp -e 300","metadata":{"execution":{"iopub.status.busy":"2024-04-14T05:59:14.602724Z","iopub.execute_input":"2024-04-14T05:59:14.603002Z","iopub.status.idle":"2024-04-14T06:00:07.154936Z","shell.execute_reply.started":"2024-04-14T05:59:14.602979Z","shell.execute_reply":"2024-04-14T06:00:07.153748Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%cd /kaggle/working","metadata":{"execution":{"iopub.status.busy":"2024-04-14T06:00:07.156823Z","iopub.execute_input":"2024-04-14T06:00:07.157723Z","iopub.status.idle":"2024-04-14T06:00:07.164471Z","shell.execute_reply.started":"2024-04-14T06:00:07.157681Z","shell.execute_reply":"2024-04-14T06:00:07.163643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls /kaggle/working/hms/results/mlp_config_0/fold_0.ckpt\n!ls /kaggle/working/hms/results/mlp_config_1/fold_0.ckpt","metadata":{"execution":{"iopub.status.busy":"2024-04-14T06:00:07.175294Z","iopub.execute_input":"2024-04-14T06:00:07.175542Z","iopub.status.idle":"2024-04-14T06:00:09.236512Z","shell.execute_reply.started":"2024-04-14T06:00:07.175516Z","shell.execute_reply":"2024-04-14T06:00:09.235311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = []\nconfigs = [\n    'newdata_hms_chris_fmax30_50sec_16ims_bandpass_9ep',\n    'newdata_hms_chris_fmax60_40sec_16ims_bandpass_9ep', # pretrain twice\n    'newdata_hms_chris_fmax30_30sec_8ims_bandpass_spe_and_eeg_9ep',\n    'newdata_hms_chris_fmax90_10sec_8ims_bandpass_spe_and_eeg_9ep',\n]\n\ntrue_cols = ['seizure','lpd','gpd','lrda','grda','other']\nfor c in configs:\n    pred_df = pd.read_csv(f'pred_{c}.csv').sort_values('eeg_id')\n    pred_cols = [f'{c}_pred_{col}' for col in true_cols]\n    preds.append(pred_df[pred_cols].values)\n\ncols = ['seizure_vote','lpd_vote','gpd_vote','lrda_vote','grda_vote','other_vote']\nadd_configs = [\n#  'tiny_vit_21m_384_el30_mixup_50ep',\n#  'caformer_s18_2_5d_256_el30_mixup_100ep',\n#  'convnext_large_384_el40_mixup_50ep',\n\n#  'convnext_large_384_el40_mixup_50ep',\n 'caformer_s18_2_5d_256_el30_mixup_100ep',\n 'tiny_vit_21m_512_el30_mixup_50ep',                \n    \n#  'eeg022_16ep_sc03c',\n 'v5_eeg_24ep_cutmix',\n 'yu4u',\n# \"1d_old_split_wd1e-2_stem_bs192_syncbn_dp0.2_ep128\",    \n]\nfor c in add_configs:\n    sub = pd.read_csv(f'{c}.csv').sort_values('eeg_id')\n    preds.append(sub[cols].values)\n\nconfigs = configs+add_configs\nfor c, pr in zip(configs, preds):\n    sub[[f'pred_{col}_{c}' for col in cols]] = pr\nsub    ","metadata":{"execution":{"iopub.status.busy":"2024-04-14T06:00:09.238124Z","iopub.execute_input":"2024-04-14T06:00:09.238455Z","iopub.status.idle":"2024-04-14T06:00:09.315817Z","shell.execute_reply.started":"2024-04-14T06:00:09.238427Z","shell.execute_reply":"2024-04-14T06:00:09.314955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MultiLayerPerceptronBase(nn.Module):\n    def __init__(self, num_classes, input_num):\n        super(MultiLayerPerceptronBase, self).__init__()\n        self.mlp = nn.Sequential(\n            nn.Linear(input_num, num_classes),\n        )\n\n    def forward(self, x, labels=None):\n        x = self.mlp(x)\n        return x\nclass MultiLayerPerceptronBaseNoBias(nn.Module):\n    def __init__(self, num_classes, input_num):\n        super(MultiLayerPerceptronBaseNoBias, self).__init__()\n        self.mlp = nn.Sequential(\n            nn.Linear(input_num, num_classes, bias=False),\n        )\n\n    def forward(self, x, labels=None):\n        x = self.mlp(x)\n        return x\nclass MultiLayerPerceptronBase2(nn.Module):\n    def __init__(self, num_classes, input_num):\n        super(MultiLayerPerceptronBase2, self).__init__()\n        self.seizure = nn.Linear(input_num//6, 1)\n        self.lpd = nn.Linear(input_num//6, 1)\n        self.gpd = nn.Linear(input_num//6, 1)\n        self.lrda = nn.Linear(input_num//6, 1)\n        self.grda = nn.Linear(input_num//6, 1)\n        self.other = nn.Linear(input_num//6, 1)\n        self.mlps = [\n            self.seizure,\n            self.lpd,\n            self.gpd,\n            self.lrda,\n            self.grda,\n            self.other,\n        ]\n    def forward(self, x, labels=None):\n        xs = []\n        for i, mlp in enumerate(self.mlps):\n            idxes = [j for j in range(x.size(1)) if j % 6 == i]\n            idxes = torch.tensor(idxes).to(x.device)\n            class_x = x[:, idxes]\n            class_x = mlp(class_x)\n            xs.append(class_x)\n\n        xs = torch.cat(xs, dim=1)\n        return xs    \nclass MultiLayerPerceptronBaseNoBias2(nn.Module):\n    def __init__(self, num_classes, input_num):\n        super(MultiLayerPerceptronBaseNoBias2, self).__init__()\n        self.seizure = nn.Linear(input_num//6, 1, bias=False)\n        self.lpd = nn.Linear(input_num//6, 1, bias=False)\n        self.gpd = nn.Linear(input_num//6, 1, bias=False)\n        self.lrda = nn.Linear(input_num//6, 1, bias=False)\n        self.grda = nn.Linear(input_num//6, 1, bias=False)\n        self.other = nn.Linear(input_num//6, 1, bias=False)\n        self.mlps = [\n            self.seizure,\n            self.lpd,\n            self.gpd,\n            self.lrda,\n            self.grda,\n            self.other,\n        ]\n    def forward(self, x, labels=None):\n        xs = []\n        for i, mlp in enumerate(self.mlps):\n            idxes = [j % 6 == i for j in range(x.size(1))]\n            idxes = torch.tensor(idxes).to(x.device)\n            class_x = x[:, idxes]\n            class_x = mlp(class_x)\n            xs.append(class_x)\n\n        xs = torch.cat(xs, dim=1)\n        return xs    \n","metadata":{"execution":{"iopub.status.busy":"2024-04-14T06:00:09.317160Z","iopub.execute_input":"2024-04-14T06:00:09.317478Z","iopub.status.idle":"2024-04-14T06:00:09.335350Z","shell.execute_reply.started":"2024-04-14T06:00:09.317452Z","shell.execute_reply":"2024-04-14T06:00:09.334496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MLPDataset(Dataset):\n    def __init__(self, cfg):\n        self.metas = cfg.df[cfg.meta_cols].values\n\n    def __len__(self):\n        return len(self.metas)\n\n    def __getitem__(self, idx):\n        meta = torch.FloatTensor(self.metas[idx])\n        return meta","metadata":{"execution":{"iopub.status.busy":"2024-04-14T06:00:09.336568Z","iopub.execute_input":"2024-04-14T06:00:09.336894Z","iopub.status.idle":"2024-04-14T06:00:09.350148Z","shell.execute_reply.started":"2024-04-14T06:00:09.336862Z","shell.execute_reply":"2024-04-14T06:00:09.349295Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class stacking_cfg0():\n    def __init__(self, df):\n        self.df = df\n        self.batch_size=4\n        self.label_features = ['seizure_vote','lpd_vote','gpd_vote','lrda_vote','grda_vote','other_vote']\n        self.num_classes = len(self.label_features)\n\n        self.meta_cols = []\n        for config in configs:\n            for col in self.label_features:\n                self.meta_cols.append(f'pred_{col}_{config}')\n\n        self.model_paths = ['/kaggle/working/hms/results/mlp_config_0/fold_0.ckpt']\n\nclass stacking_cfg1(stacking_cfg0):\n    def __init__(self, df):\n        super().__init__(df)\n        self.model_paths = ['/kaggle/working/hms/results/mlp_config_1/fold_0.ckpt']","metadata":{"execution":{"iopub.status.busy":"2024-04-14T06:00:09.351291Z","iopub.execute_input":"2024-04-14T06:00:09.351618Z","iopub.status.idle":"2024-04-14T06:00:09.361988Z","shell.execute_reply.started":"2024-04-14T06:00:09.351584Z","shell.execute_reply":"2024-04-14T06:00:09.361269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models = []\n\ncfg = stacking_cfg0(sub.copy())\nfor model_path in cfg.model_paths:\n    model = MultiLayerPerceptronBase(num_classes=cfg.num_classes, input_num=len(cfg.meta_cols))\n    print(model_path)\n    m = {}\n    state_dict = torch.load(model_path, map_location=torch.device('cpu'))['state_dict']\n    for k, v in dict(state_dict).items():\n        m[k[6:]] = v    \n    model.load_state_dict(m)\n    model.to(DEVICE)\n    model.eval()\n    models.append(model)\n\ncfg = stacking_cfg1(sub.copy())\nfor model_path in cfg.model_paths:\n    model = MultiLayerPerceptronBaseNoBias(num_classes=cfg.num_classes, input_num=len(cfg.meta_cols))\n    print(model_path)\n    m = {}\n    state_dict = torch.load(model_path, map_location=torch.device('cpu'))['state_dict']\n    for k, v in dict(state_dict).items():\n        m[k[6:]] = v    \n    model.load_state_dict(m)\n    model.to(DEVICE)\n    model.eval()\n    models.append(model)\n    ","metadata":{"execution":{"iopub.status.busy":"2024-04-14T06:00:09.363134Z","iopub.execute_input":"2024-04-14T06:00:09.363415Z","iopub.status.idle":"2024-04-14T06:00:09.380908Z","shell.execute_reply.started":"2024-04-14T06:00:09.363392Z","shell.execute_reply":"2024-04-14T06:00:09.380106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = MLPDataset(cfg)\nloader = DataLoader(ds, batch_size=cfg.batch_size, shuffle=False, drop_last=False, num_workers=4)\npreds = []\nfor 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    preds += np.mean(batch_preds, axis=0).tolist()\nname = cfg.model_paths[0].split('/')[3].replace('-', '_')\ncfg.pre = name\ncfg.df[cfg.label_features] = softmax(np.array(preds), axis=1)\nsub = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/sample_submission.csv')","metadata":{"execution":{"iopub.status.busy":"2024-04-14T06:00:09.382021Z","iopub.execute_input":"2024-04-14T06:00:09.382372Z","iopub.status.idle":"2024-04-14T06:00:09.814112Z","shell.execute_reply.started":"2024-04-14T06:00:09.382329Z","shell.execute_reply":"2024-04-14T06:00:09.812750Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -rf ./*","metadata":{"execution":{"iopub.status.busy":"2024-04-14T06:00:09.815987Z","iopub.execute_input":"2024-04-14T06:00:09.816399Z","iopub.status.idle":"2024-04-14T06:00:10.913622Z","shell.execute_reply.started":"2024-04-14T06:00:09.816363Z","shell.execute_reply":"2024-04-14T06:00:10.912352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cfg.df[list(sub)].to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2024-04-14T06:00:10.915104Z","iopub.execute_input":"2024-04-14T06:00:10.915435Z","iopub.status.idle":"2024-04-14T06:00:10.924736Z","shell.execute_reply.started":"2024-04-14T06:00:10.915405Z","shell.execute_reply":"2024-04-14T06:00:10.923975Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}