{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":59093,"databundleVersionId":7469972},{"sourceType":"datasetVersion","sourceId":7754261,"datasetId":4532886,"databundleVersionId":7854993},{"sourceType":"datasetVersion","sourceId":7392775,"datasetId":4297782,"databundleVersionId":7483780},{"sourceType":"datasetVersion","sourceId":7826851,"datasetId":4586546,"databundleVersionId":7930464},{"sourceType":"datasetVersion","sourceId":7392733,"datasetId":4297749,"databundleVersionId":7483738},{"sourceType":"datasetVersion","sourceId":7652061,"datasetId":4216847,"databundleVersionId":7748850},{"sourceType":"datasetVersion","sourceId":7994182,"datasetId":4607396,"databundleVersionId":8105428},{"sourceType":"kernelVersion","sourceId":164443259},{"sourceType":"kernelVersion","sourceId":169973997},{"sourceType":"kernelVersion","sourceId":169994548}],"dockerImageVersionId":30648,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Directory settings","metadata":{}},{"cell_type":"markdown","source":"# About this notebook\n\nApply knowledge distillation on stage1 trained models\n\n## Version 1\nSame as my [baseline](https://www.kaggle.com/code/medali1992/hms-efficientnetb0-train) with few modification:\n* Changed the model name to `tf_efficientnet_b0.ns_jft_in1k`\n* Changed the CV scheme\n* https://www.kaggle.com/code/medali1992/hms-efficientnetb0-strat-train\n\n## Version 2\n* Replace contrastive loss with mse_loss","metadata":{}},{"cell_type":"code","source":"!pip install /kaggle/input/wheel-albumentation/albumentations-1.4.0-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:42:44.008059Z","iopub.execute_input":"2024-04-03T03:42:44.008487Z","iopub.status.idle":"2024-04-03T03:42:58.280938Z","shell.execute_reply.started":"2024-04-03T03:42:44.008456Z","shell.execute_reply":"2024-04-03T03:42:58.279891Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# directory settings\n# ====================================================\n\nimport os\n\nOUTPUT_DIR = './'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)\n    \nPOP_2_DIR = OUTPUT_DIR + 'pop_2_weight_oof/'\nif not os.path.exists(POP_2_DIR):\n    os.makedirs(POP_2_DIR)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-04-03T03:42:58.283464Z","iopub.execute_input":"2024-04-03T03:42:58.284010Z","iopub.status.idle":"2024-04-03T03:42:58.290623Z","shell.execute_reply.started":"2024-04-03T03:42:58.283966Z","shell.execute_reply":"2024-04-03T03:42:58.289652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Libraries","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nfrom glob import glob\nimport sys\nimport math\nimport time\nimport random\nimport shutil\nfrom pathlib import Path\nfrom typing import Dict, List\nfrom scipy.stats import entropy\nfrom scipy.signal import butter, lfilter, freqz\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\nsys.path.append('/kaggle/input/kaggle-kl-div')\nfrom kaggle_kl_div import score\nimport numpy as np\nimport pandas as pd\nfrom sklearn import preprocessing\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.metrics import accuracy_score, log_loss\nfrom tqdm.auto import tqdm\nfrom functools import partial\nimport cv2\nfrom PIL import Image\nimport torch\nimport torch.nn as nn\nimport pytorch_lightning as pl\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD, AdamW\nimport torchvision.models as models\nfrom torch.nn.parameter import Parameter\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau, OneCycleLR, CosineAnnealingLR, CosineAnnealingWarmRestarts\nfrom sklearn.preprocessing import LabelEncoder\nfrom torchvision.transforms import v2\nfrom sklearn.model_selection import GroupKFold\nfrom sklearn.model_selection import train_test_split\nimport albumentations as A\nfrom albumentations import (Compose, Normalize, Resize, RandomResizedCrop, HorizontalFlip, VerticalFlip, ShiftScaleRotate, Transpose)\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\nimport timm\nimport warnings \nwarnings.filterwarnings('ignore')\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nfrom matplotlib import pyplot as plt\nimport joblib\nos.environ['CUDA_VISIBLE_DEVICES'] = \"0,1\"\nVERSION=3","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:42:58.291964Z","iopub.execute_input":"2024-04-03T03:42:58.292566Z","iopub.status.idle":"2024-04-03T03:43:10.756609Z","shell.execute_reply.started":"2024-04-03T03:42:58.292534Z","shell.execute_reply":"2024-04-03T03:43:10.755791Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# CFG\n# ====================================================\n\nclass CFG:\n    wandb = False\n    debug = False\n    train=True\n    apex=True\n    stage1_pop1=False\n    stage2_pop2=True\n    VISUALIZE=True\n    FREEZE=False\n    SparK=False\n    t4_gpu=True\n    scheduler='OneCycleLR' # ['ReduceLROnPlateau', 'CosineAnnealingLR', 'CosineAnnealingWarmRestarts','OneCycleLR']\n    # CosineAnnealingLR params\n    cosanneal_params={\n        'T_max':6,\n        'eta_min':1e-5,\n        'last_epoch':-1\n    }\n    #ReduceLROnPlateau params\n    reduce_params={\n        'mode':'min',\n        'factor':0.2,\n        'patience':4,\n        'eps':1e-6,\n        'verbose':True\n    }\n    # CosineAnnealingWarmRestarts params\n    cosanneal_res_params={\n        'T_0':20,\n        'eta_min':1e-6,\n        'T_mult':1,\n        'last_epoch':-1\n    }\n    print_freq=30\n    num_workers = 1\n    model = 'tf_efficientnet_b0.ns_jft_in1k'\n    model_name = 'stage2_knowledge_distillation_tf_effb0_ns'\n    optimizer='AdamW'\n    epochs = 5\n    factor = 0.9\n    patience = 2\n    eps = 1e-6\n    lr = 1e-3\n    MIXUP_CUTMIX_PROB = 0.5\n    SPEC_SIZE  = (512, 512, 3)\n    min_lr = 1e-6\n    batch_size = 64\n    weights=[0.5, 0.5, 1]\n    weight_decay = 1e-2\n    batch_scheduler=True\n    gradient_accumulation_steps = 1\n    max_grad_norm = 1e7\n    seed = 2024\n    target_cols = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\n    target_size = 7\n    pred_cols = ['pred_seizure_vote', 'pred_lpd_vote', 'pred_gpd_vote', 'pred_lrda_vote', 'pred_grda_vote', 'pred_other_vote']\n    n_fold = 5\n    trn_fold = [0, 1, 2, 3, 4]\n    PATH = '/kaggle/input/hms-harmful-brain-activity-classification/'\n    data_root = \"/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/\"\n    preproc_specs_path = \"/kaggle/input/hms-final-zimmermann-specs/zimmermann_images/\"\n    ","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:43:10.757774Z","iopub.execute_input":"2024-04-03T03:43:10.758044Z","iopub.status.idle":"2024-04-03T03:43:10.768371Z","shell.execute_reply.started":"2024-04-03T03:43:10.758020Z","shell.execute_reply":"2024-04-03T03:43:10.767508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"def init_logger(log_file=OUTPUT_DIR+'train.log'):\n    from logging import getLogger, INFO, FileHandler,  Formatter,  StreamHandler\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\nLOGGER = init_logger()\n\ndef get_score(preds, targets):\n    oof = pd.DataFrame(preds.copy())\n    oof['id'] = np.arange(len(oof))\n\n    true = pd.DataFrame(targets.copy())\n    true['id'] = np.arange(len(true))\n\n    cv = score(solution=true, submission=oof, row_id_column_name='id')\n    return cv\n\ndef butter_bandpass(lowcut, highcut, fs, order=5):\n    return butter(order, [lowcut, highcut], fs=fs, btype='band')\n\ndef butter_bandpass_filter(data, lowcut, highcut, fs, order=5):\n    b, a = butter_bandpass(lowcut, highcut, fs, order=order)\n    y = lfilter(b, a, data)\n    return y\n\n\ndef denoise_filter(x):\n    # Sample rate and desired cutoff frequencies (in Hz).\n    fs = 200.0\n    lowcut = 1.0\n    highcut = 25.0\n    \n    # Filter a noisy signal.\n    T = 50\n    nsamples = T * fs\n    t = np.arange(0, nsamples) / fs\n    y = butter_bandpass_filter(x, lowcut, highcut, fs, order=6)\n    y = (y + np.roll(y,-1)+ np.roll(y,-2)+ np.roll(y,-3))/4\n    y = y[0:-1:4]\n    \n    return y\n\nclass KLDivLossWithLogits(nn.KLDivLoss):\n\n    def __init__(self):\n        super().__init__(reduction=\"batchmean\")\n\n    def forward(self, y, t):\n        y = nn.functional.log_softmax(y,  dim=1)\n        loss = super().forward(y, t)\n\n        return loss\n\nclass ContrastiveLoss(torch.nn.Module):\n    \"\"\"\n    Contrastive loss function.\n    Based on: http://yann.lecun.com/exdb/publis/pdf/hadsell-chopra-lecun-06.pdf\n    \"\"\"\n\n    def __init__(self, margin=1.0):\n        super(ContrastiveLoss, self).__init__()\n        self.margin = margin\n\n    def forward(self, output1, output2, label):\n        euclidean_distance = F.cosine_similarity(F.normalize(output1), F.normalize(output2))\n        loss_contrastive = torch.mean((1-label) * torch.pow(euclidean_distance, 2) +\n                                      (label) * torch.pow(torch.clamp(self.margin - euclidean_distance, min=0.0), 2))\n\n\n        return loss_contrastive\n\ndef seed_torch(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.enabled =False\n    \nseed_torch(seed=CFG.seed)","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:43:10.771072Z","iopub.execute_input":"2024-04-03T03:43:10.771368Z","iopub.status.idle":"2024-04-03T03:43:10.794970Z","shell.execute_reply.started":"2024-04-03T03:43:10.771345Z","shell.execute_reply":"2024-04-03T03:43:10.793761Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load train data","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/train.csv')\nlabel_cols = df.columns[-6:]\ndf['total_evaluators'] = df[['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']].sum(axis=1)\n\nprint(f\"Train cataframe shape is: {df.shape}\")\nprint(f\"Labels: {list(label_cols)}\")\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:43:10.795974Z","iopub.execute_input":"2024-04-03T03:43:10.796258Z","iopub.status.idle":"2024-04-03T03:43:11.100492Z","shell.execute_reply.started":"2024-04-03T03:43:10.796235Z","shell.execute_reply":"2024-04-03T03:43:11.099524Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Deduplicate Train EEG Id","metadata":{}},{"cell_type":"code","source":"train = df.groupby('eeg_id')[['spectrogram_id','spectrogram_label_offset_seconds', 'eeg_label_offset_seconds', 'total_evaluators']].agg({\n    'spectrogram_id':'first',\n    'total_evaluators':'first',\n    'spectrogram_label_offset_seconds':'min',\n    'eeg_label_offset_seconds':'min'\n    \n})\ntrain.columns = ['spectrogram_id', 'total_evaluators', 'spectrogram_label_offset_seconds_min', 'eeg_label_offset_seconds_min']\n\naux = df.groupby('eeg_id')[['spectrogram_id','spectrogram_label_offset_seconds']].agg({\n    'spectrogram_label_offset_seconds':'max'\n})\ntrain['spectrogram_label_offset_seconds_max'] = aux\n\naux = df.groupby('eeg_id')[['patient_id']].agg('first')\ntrain['patient_id'] = aux\n\naux = df.groupby('eeg_id')[label_cols].agg('sum')\nfor label in label_cols:\n    train[label] = aux[label].values\n    \ny_data = train[label_cols].values\ny_data = y_data / y_data.sum(axis=1,keepdims=True)\ntrain[label_cols] = y_data\n\naux = df.groupby('eeg_id')[['expert_consensus']].agg('first')\ntrain['target'] = aux\n\ntrain = train.reset_index()\ntrain['aux_label'] = train['total_evaluators'].map(lambda x: 1.0 if x < 10 else 0.0)\n\nprint('Train non-overlapp eeg_id shape:', train.shape )\ntrain.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:43:11.102174Z","iopub.execute_input":"2024-04-03T03:43:11.102888Z","iopub.status.idle":"2024-04-03T03:43:11.209258Z","shell.execute_reply.started":"2024-04-03T03:43:11.102848Z","shell.execute_reply":"2024-04-03T03:43:11.208260Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CV Scheme","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import StratifiedGroupKFold\nimport numpy as np\n\nbin_edges = [0, 5, 10, 15, 20, np.inf]\nnum_bins = len(bin_edges) - 1\ntrain['evaluator_bin'] = pd.cut(train['total_evaluators'], bins=bin_edges, labels=False)\npatient_groups = train.groupby('patient_id').ngroup()\nstratified_kfold = StratifiedGroupKFold(n_splits=CFG.n_fold)\ntrain_indices = []\nvalid_indices = []\nfor fold, (train_idx, valid_idx) in enumerate(stratified_kfold.split(X=train, y=train['evaluator_bin'], groups=patient_groups)):\n    train_indices.append(train_idx)\n    valid_indices.append(valid_idx) \n    train.loc[valid_idx, \"fold\"] = fold\n\nplt.figure(figsize=(8, 6))\nplt.hist(train['evaluator_bin'], bins=num_bins, color='green', edgecolor='black', align='left')\nplt.title('Distribution of Stratified Bins')\nplt.xlabel('Evaluator Bin')\nplt.ylabel('Frequency')\nplt.grid(True)\nplt.xticks(range(num_bins), [f'{bin_edges[i]}-{bin_edges[i+1]}' for i in range(num_bins)])\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:43:11.210354Z","iopub.execute_input":"2024-04-03T03:43:11.210748Z","iopub.status.idle":"2024-04-03T03:43:12.414553Z","shell.execute_reply.started":"2024-04-03T03:43:11.210699Z","shell.execute_reply":"2024-04-03T03:43:12.413685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.fold.value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:43:12.415552Z","iopub.execute_input":"2024-04-03T03:43:12.415834Z","iopub.status.idle":"2024-04-03T03:43:12.427829Z","shell.execute_reply.started":"2024-04-03T03:43:12.415810Z","shell.execute_reply":"2024-04-03T03:43:12.426970Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import seaborn as sns\nplt.figure(figsize=(8, 6))\nsns.countplot(data=train, x='fold')\nplt.title('Distribution of Fold IDs')\nplt.xlabel('Fold ID')\nplt.ylabel('Count')\nplt.grid(True)\nplt.show()\n# Plot distribution of evaluator bins within folds\nplt.figure(figsize=(8, 6))\nsns.countplot(data=train, x='fold', hue='evaluator_bin')\nplt.title('Distribution of Evaluator Bins within Folds')\nplt.xlabel('Fold ID')\nplt.ylabel('Count')\nplt.legend(title='Evaluator Bin', loc='upper right')\nplt.grid(True)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:43:12.429171Z","iopub.execute_input":"2024-04-03T03:43:12.429506Z","iopub.status.idle":"2024-04-03T03:43:13.010679Z","shell.execute_reply.started":"2024-04-03T03:43:12.429473Z","shell.execute_reply":"2024-04-03T03:43:13.009793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# W&B Settings","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# wandb\n# ====================================================\nif CFG.wandb:\n    \n    import wandb\n\n    try:\n        from kaggle_secrets import UserSecretsClient\n        user_secrets = UserSecretsClient()\n        secret_value_0 = user_secrets.get_secret(\"wandb_key\")\n        wandb.login(key=secret_value_0)\n        anony = None\n    except:\n        anony = \"must\"\n        print('If you want to use your W&B account, go to Add-ons -> Secrets and provide your W&B access token. Use the Label name as wandb_api. \\nGet your W&B access token from here: https://wandb.ai/authorize')\n\n\n    def class2dict(f):\n        return dict((name, getattr(f, name)) for name in dir(f) if not name.startswith('__'))","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:43:13.012058Z","iopub.execute_input":"2024-04-03T03:43:13.012768Z","iopub.status.idle":"2024-04-03T03:43:13.019656Z","shell.execute_reply.started":"2024-04-03T03:43:13.012732Z","shell.execute_reply":"2024-04-03T03:43:13.018742Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Mixup Augmentation","metadata":{}},{"cell_type":"code","source":"from albumentations.core.transforms_interface import ImageOnlyTransform\n\ndef mixup_data(\n        X: torch.Tensor, y: torch.Tensor, use_mixup: bool, device,\n        alpha: float = 1.0\n    ):\n    \"\"\"\n    Performs MixUp augmentation.\n    :param use_mixup: whether to use mixup or not.\n    :param X: batch with images.\n    :param y: ground truth.\n    :param alpha: parameter to use in the beta distribution.\n    :param device: indicates if using CPU or GPU.\n    :return mixed_X: mixed X.\n    :return y_a: class of the original image.\n    :return y_b: class of the image used for mixing.\n    :return lambda: lamdba parameter that indicates the percentage of the mix.\n    \"\"\"\n    if not use_mixup:\n        return X, y, None, None\n\n    if alpha > 0:\n        lam = np.random.beta(alpha, alpha) # draw random number from beta distribution\n    else:\n        lam = 1\n\n    batch_size = X.size()[0]\n    index = torch.randperm(batch_size).to(device) # torch tensor with shuffled numbers between 1:batch_size\n    mixed_X = lam * X + (1 - lam) * X[index, :] # perform mixup to the whole batch\n    y_a, y_b = y, y[index]\n    return mixed_X, y_a, y_b, lam\n\ndef rand_bbox(size, lam):\n    W = size[2]\n    H = size[3]\n    cut_rat = np.sqrt(1. - lam)\n    cut_w = np.int_(W * cut_rat)\n    cut_h = np.int_(H * cut_rat)\n\n    # uniform\n    cx = np.random.randint(W)\n    cy = np.random.randint(H)\n\n    bbx1 = np.clip(cx - cut_w // 2, 0, W)\n    bby1 = np.clip(cy - cut_h // 2, 0, H)\n    bbx2 = np.clip(cx + cut_w // 2, 0, W)\n    bby2 = np.clip(cy + cut_h // 2, 0, H)\n\n    return bbx1, bby1, bbx2, bby2\n\ndef cutmix(X: torch.Tensor, y: torch.Tensor, use_cutmix: bool, device, alpha: float = 1.0):\n    \n    if not use_cutmix:\n        return X, y, None, None\n    \n    lam = np.random.beta(alpha, alpha)\n    rand_index = torch.randperm(X.size()[0]).to(device)\n    target_a = y\n    target_b = y[rand_index]\n    bbx1, bby1, bbx2, bby2 = rand_bbox(X.size(), lam)\n    X[:, :, bbx1:bbx2, bby1:bby2] = X[rand_index, :, bbx1:bbx2, bby1:bby2]\n    # adjust lambda to exactly match pixel ratio\n    lam = 1 - ((bbx2 - bbx1) * (bby2 - bby1) / (X.size()[-1] * X.size()[-2]))\n    return X, target_a, target_b, lam\n\n\ndef get_criterion(use_mixup_cutmix, criterion):\n    \"\"\"\n    This function computes the criterion/loss depending whether MixUp augmentation was applied or not.\n    If MixUp was applied it returns a weighted average of the loss averaging by the lambda parameter.\n    Otherwise, it returns the regular loss as MixUp was not applied.\n    :param config: configuration class with param to use mixup or not.\n    :param criterion: loss function to use.\n    \"\"\"\n\n    def mixup_criterion(pred, y_a, y_b, lam):\n        return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)\n\n    def single_criterion(pred, y_a, y_b, lam):\n        return criterion(pred, y_a)\n\n    if use_mixup_cutmix:\n        return mixup_criterion\n    else:\n        return single_criterion\n\n\nclass RandomShuffleChannels(ImageOnlyTransform):\n    \"\"\"\n    Swaps channels on an 8-channel array according to different strategies.\n    :param p: probability to apply the transformation.\n    :param which: the strategy to swap channels: \"first\" randomly swaps the first 4 channels,\n    \"last\" swaps the last 4 channels and \"both\" swaps all channels.\n    \"\"\"\n    def __init__(self, safe_db_lists=[], p: float = 0.5, which: str = \"both\") -> None:\n        super(RandomShuffleChannels, self).__init__()\n        self.safe_db_lists = safe_db_lists\n        self.p = p\n        self.which = which\n\n    def apply(self, img, copy=True, **params):\n        if np.random.uniform(0, 1) > self.p:\n            return img\n        if copy:\n            img = img.copy()\n\n        if self.which == \"first\":\n            permuted_indices = np.random.permutation([0,1,2,3])\n            indices = np.concatenate([permuted_indices, np.array([4,5,6,7])])\n            swapped_image = img[:, :, indices]\n        elif self.which == \"last\":\n            permuted_indices = np.random.permutation([4,5,6,7])\n            indices = np.concatenate([np.array([0,1,2,3]), permuted_indices])\n            swapped_image = img[:, :, indices]\n        else:\n            permuted_indices_a = np.random.permutation([0,1,2,3])\n            permuted_indices_b = np.random.permutation([4,5,6,7])\n            indices = np.concatenate([permuted_indices_a, permuted_indices_b])\n            swapped_image = img[:, :, indices]\n        return swapped_image","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:43:13.021137Z","iopub.execute_input":"2024-04-03T03:43:13.021583Z","iopub.status.idle":"2024-04-03T03:43:13.046845Z","shell.execute_reply.started":"2024-04-03T03:43:13.021551Z","shell.execute_reply":"2024-04-03T03:43:13.045905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(\n        self, df: pd.DataFrame,\n        augment: bool = False, mode: str = 'train',\n    ): \n        self.df = df\n        self.augment = augment\n        self.mode = mode\n        self.mixup_prob = CFG.MIXUP_CUTMIX_PROB\n        \n    def __len__(self):\n        \"\"\"\n        Denotes the number of batches per epoch.\n        \"\"\"\n        return len(self.df)\n        \n    def __getitem__(self, index):\n        \"\"\"\n        Generate one batch of data.\n        \"\"\"\n        output_dict = {}\n        X, y = self.__data_generation(index)\n        p = np.random.uniform(0,1)\n        if self.augment:\n            X = self.__transform(X) \n            output_dict[\"spectrogram\"] = torch.tensor(X, dtype=torch.float32)\n            output_dict[\"labels\"] = torch.tensor(y, dtype=torch.float32)\n            return output_dict\n        else:\n            output_dict[\"spectrogram\"] = torch.tensor(X, dtype=torch.float32)\n            output_dict[\"labels\"] = torch.tensor(y, dtype=torch.float32)\n            return output_dict\n                        \n    def __data_generation(self, index):\n        \"\"\"\n        Generates data containing batch_size samples.\n        \"\"\"\n        #X = np.zeros(CFG.SPEC_SIZE, dtype='float32')\n        y = np.zeros(7, dtype='float32')\n        row = self.df.iloc[index]\n        eeg_id = row['eeg_id']\n        spec_offset = int(row['spectrogram_label_offset_seconds_min'])\n        eeg_offset = int(row['eeg_label_offset_seconds_min'])\n        file_path = f'/kaggle/input/3-diff-time-specs-hms/images/{eeg_id}_{spec_offset}_{eeg_offset}.npz'\n        data = np.load(file_path)\n        eeg_data = data['final_image']\n        eeg_data_expanded = np.repeat(eeg_data[:, :, np.newaxis], 3, axis=2)\n        #X = eeg_data_expanded\n        \n        if self.mode != 'test':\n            \n            y[:-1] = row[CFG.target_cols]\n            y[-1] = row['aux_label']\n            \n        return eeg_data_expanded, y\n    \n    def __transform(self, img):\n        params1 = {\n                    \"num_masks_x\": 1,    \n                    \"mask_x_length\": (1, 10), # This line changed from fixed  to a range\n                    \"fill_value\": (0, 1, 2),\n                    }\n        params2 = {    \n                    \"num_masks_y\": 1,    \n                    \"mask_y_length\": (1, 10),\n                    \"fill_value\": (0, 1, 2),    \n                    }\n        params3 = {    \n                    \"num_masks_x\": (2, 4),\n                    \"num_masks_y\": 5,    \n                    \"mask_y_length\": (1, 10),\n                    \"mask_x_length\": (1, 10),\n                    \"fill_value\": (0, 1, 2),  \n                    }\n        \n        transforms = A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.XYMasking(**params1, p=0.3),\n            A.XYMasking(**params2, p=0.3),\n            A.XYMasking(**params3, p=0.3),\n        ])\n        return transforms(image=img)['image']","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:43:13.048275Z","iopub.execute_input":"2024-04-03T03:43:13.048576Z","iopub.status.idle":"2024-04-03T03:43:13.064258Z","shell.execute_reply.started":"2024-04-03T03:43:13.048552Z","shell.execute_reply":"2024-04-03T03:43:13.063354Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataloader","metadata":{}},{"cell_type":"code","source":"dataset = CustomDataset(train, augment=True, mode=\"train\")\ndataloader = DataLoader(dataset, batch_size=32, shuffle=False)\n\nbatch = dataset[0]\nX, y = batch[\"spectrogram\"], batch[\"labels\"]\nprint(f\"X shape: {X.shape}\")\nprint(f\"y shape: {y.shape}\")\n\ndel dataset, X, y\n_ = gc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:43:13.069473Z","iopub.execute_input":"2024-04-03T03:43:13.069769Z","iopub.status.idle":"2024-04-03T03:43:13.350543Z","shell.execute_reply.started":"2024-04-03T03:43:13.069743Z","shell.execute_reply":"2024-04-03T03:43:13.349450Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.VISUALIZE:\n    ROWS = 2\n    COLS = 3\n    for batch in dataloader:\n        X, y = batch[\"spectrogram\"], batch[\"labels\"]\n        plt.figure(figsize=(20,8))\n        plt.imshow(X[0].numpy(), cmap='jet')\n        plt.axis('off')\n        break\n        \ndel dataloader\n_ = gc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:43:13.351592Z","iopub.execute_input":"2024-04-03T03:43:13.351905Z","iopub.status.idle":"2024-04-03T03:43:14.731313Z","shell.execute_reply.started":"2024-04-03T03:43:13.351881Z","shell.execute_reply":"2024-04-03T03:43:14.730355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class AdaptiveGeM(nn.Module):\n        def __init__(self, size=(1, 1), p=3, eps=1e-6):\n            super().__init__()\n            self.size = size\n            self.p = Parameter(torch.ones(1)*p)\n            self.eps = eps\n\n        def forward(self, x):\n            return F.adaptive_avg_pool2d(\n                x.clamp(min=self.eps).pow(self.p), self.size).pow(1./self.p)\n\n        def __repr__(self):\n            return f'AdaptiveGeM(size={self.size}, p={self.p}, eps={self.eps})'\n\nclass CustomModel(nn.Module):\n    def __init__(self, config, num_classes: int = 7, pretrained: bool = True):\n        super(CustomModel, self).__init__()\n        self.USE_KAGGLE_SPECTROGRAMS = False\n        self.USE_EEG_SPECTROGRAMS = True\n        self.model = timm.create_model(\n            config.model,\n            pretrained=pretrained,\n        )\n        if config.FREEZE:\n            for i,(name, param) in enumerate(list(self.model.named_parameters())\\\n                                             [0:config.NUM_FROZEN_LAYERS]):\n                param.requires_grad = False\n\n        self.features = nn.Sequential(*list(self.model.children())[:-2])\n        \n        self.custom_layers = nn.Sequential(\n            AdaptiveGeM(1),\n            nn.Flatten(),\n            nn.Linear(self.model.num_features, num_classes)\n        )\n\n    def __reshape_input(self, x):\n        \"\"\"\n        Reshapes input (128, 256, 8) -> (512, 512, 3) monotone image.\n        \"\"\" \n        # === Get spectograms ===\n        spectograms = [x[:, :, :, i:i+1] for i in range(4)]\n        spectograms = torch.cat(spectograms, dim=1)\n        \n        # === Get EEG spectograms ===\n        eegs = [x[:, :, :, i:i+1] for i in range(4,8)]\n        eegs = torch.cat(eegs, dim=1)\n        \n        # === Reshape (512,512,3) ===\n        if self.USE_KAGGLE_SPECTROGRAMS & self.USE_EEG_SPECTROGRAMS:\n            x = torch.cat([spectograms, eegs], dim=2)\n        elif self.USE_EEG_SPECTROGRAMS:\n            x = eegs\n        else:\n            x = spectograms\n            \n        x = torch.cat([x,x,x], dim=3)\n        x = x.permute(0, 3, 1, 2)\n        return x\n    \n    def extract_features(self, x):\n        x = self.features(x.permute(0, 3, 1, 2))\n        return x\n    \n    def forward(self, x):\n        x = self.features(x.permute(0, 3, 1, 2))\n        x = self.custom_layers(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:43:14.732822Z","iopub.execute_input":"2024-04-03T03:43:14.733107Z","iopub.status.idle":"2024-04-03T03:43:14.749772Z","shell.execute_reply.started":"2024-04-03T03:43:14.733082Z","shell.execute_reply":"2024-04-03T03:43:14.748738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomLoss(nn.Module):\n    def __init__(self, weights=[1, 1, 1]):\n        super(CustomLoss, self).__init__()\n        self.weights = weights\n        \n    def forward(self, teacher_features, student_features, teacher_pred, student_pred, labels, contrastive_target):\n        \n        contrastive_loss = ContrastiveLoss()(teacher_features, student_features, contrastive_target)\n        #mse_loss = nn.MSELoss()(teacher_features, student_features)\n        kl_loss_pred = nn.KLDivLoss(reduction=\"batchmean\")(F.log_softmax(student_pred, dim=1), labels)\n        kl_loss_teacher_student = nn.KLDivLoss(reduction=\"batchmean\")(F.log_softmax(student_pred, dim=1), F.softmax(teacher_pred, dim=1))\n        loss = self.weights[0] * contrastive_loss + self.weights[1] * kl_loss_teacher_student + kl_loss_pred * self.weights[2]\n        #loss = self.weights[0] * mse_loss + self.weights[1] * kl_loss_teacher_student + kl_loss_pred * self.weights[2]\n        return loss","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:43:14.751258Z","iopub.execute_input":"2024-04-03T03:43:14.751622Z","iopub.status.idle":"2024-04-03T03:43:14.765376Z","shell.execute_reply.started":"2024-04-03T03:43:14.751592Z","shell.execute_reply":"2024-04-03T03:43:14.764611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"iot = torch.randn(2, 512, 512, 3)\nmodel = CustomModel(CFG)\noutput = model(iot)\nprint(output.shape)\n\ndel iot, model, output\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:43:14.766388Z","iopub.execute_input":"2024-04-03T03:43:14.766658Z","iopub.status.idle":"2024-04-03T03:43:16.310117Z","shell.execute_reply.started":"2024-04-03T03:43:14.766635Z","shell.execute_reply":"2024-04-03T03:43:16.309169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Adan Optimizer","metadata":{}},{"cell_type":"code","source":"import math\nimport torch\nfrom torch.optim.optimizer import Optimizer\n\n\nclass Adan(Optimizer):\n    \"\"\"\n    Implements a pytorch variant of Adan\n    Adan was proposed in\n    Adan: Adaptive Nesterov Momentum Algorithm for Faster Optimizing Deep Models[J]. arXiv preprint arXiv:2208.06677, 2022.\n    https://arxiv.org/abs/2208.06677\n    Arguments:\n        params (iterable): iterable of parameters to optimize or dicts defining parameter groups.\n        lr (float, optional): learning rate. (default: 1e-3)\n        betas (Tuple[float, float, flot], optional): coefficients used for computing \n            running averages of gradient and its norm. (default: (0.98, 0.92, 0.99))\n        eps (float, optional): term added to the denominator to improve \n            numerical stability. (default: 1e-8)\n        weight_decay (float, optional): decoupled weight decay (L2 penalty) (default: 0)\n        max_grad_norm (float, optional): value used to clip \n            global grad norm (default: 0.0 no clip)\n        no_prox (bool): how to perform the decoupled weight decay (default: False)\n    \"\"\"\n\n    def __init__(self, params, lr=1e-3, betas=(0.98, 0.92, 0.99), eps=1e-8,\n                 weight_decay=0.2, max_grad_norm=0.0, no_prox=False):\n        if not 0.0 <= max_grad_norm:\n            raise ValueError(\"Invalid Max grad norm: {}\".format(max_grad_norm))\n        if not 0.0 <= lr:\n            raise ValueError(\"Invalid learning rate: {}\".format(lr))\n        if not 0.0 <= eps:\n            raise ValueError(\"Invalid epsilon value: {}\".format(eps))\n        if not 0.0 <= betas[0] < 1.0:\n            raise ValueError(\"Invalid beta parameter at index 0: {}\".format(betas[0]))\n        if not 0.0 <= betas[1] < 1.0:\n            raise ValueError(\"Invalid beta parameter at index 1: {}\".format(betas[1]))\n        if not 0.0 <= betas[2] < 1.0:\n            raise ValueError(\"Invalid beta parameter at index 2: {}\".format(betas[2]))\n        defaults = dict(lr=lr, betas=betas, eps=eps,\n                        weight_decay=weight_decay,\n                        max_grad_norm=max_grad_norm, no_prox=no_prox)\n        super(Adan, self).__init__(params, defaults)\n\n    def __setstate__(self, state):\n        super(Adan, self).__setstate__(state)\n        for group in self.param_groups:\n            group.setdefault('no_prox', False)\n\n    @torch.no_grad()\n    def restart_opt(self):\n        for group in self.param_groups:\n            group['step'] = 0\n            for p in group['params']:\n                if p.requires_grad:\n                    state = self.state[p]\n                    # State initialization\n\n                    # Exponential moving average of gradient values\n                    state['exp_avg'] = torch.zeros_like(p)\n                    # Exponential moving average of squared gradient values\n                    state['exp_avg_sq'] = torch.zeros_like(p)\n                    # Exponential moving average of gradient difference\n                    state['exp_avg_diff'] = torch.zeros_like(p)\n\n    @torch.no_grad()\n    def step(self):\n        \"\"\"\n            Performs a single optimization step.\n        \"\"\"\n        if self.defaults['max_grad_norm'] > 0:\n            device = self.param_groups[0]['params'][0].device\n            global_grad_norm = torch.zeros(1, device=device)\n\n            max_grad_norm = torch.tensor(self.defaults['max_grad_norm'], device=device)\n            for group in self.param_groups:\n\n                for p in group['params']:\n                    if p.grad is not None:\n                        grad = p.grad\n                        global_grad_norm.add_(grad.pow(2).sum())\n\n            global_grad_norm = torch.sqrt(global_grad_norm)\n\n            clip_global_grad_norm = torch.clamp(max_grad_norm / (global_grad_norm + group['eps']), max=1.0)\n        else:\n            clip_global_grad_norm = 1.0\n\n        for group in self.param_groups:\n            beta1, beta2, beta3 = group['betas']\n            # assume same step across group now to simplify things\n            # per parameter step can be easily support by making it tensor, or pass list into kernel\n            if 'step' in group:\n                group['step'] += 1\n            else:\n                group['step'] = 1\n\n            bias_correction1 = 1.0 - beta1 ** group['step']\n\n            bias_correction2 = 1.0 - beta2 ** group['step']\n\n            bias_correction3 = 1.0 - beta3 ** group['step']\n\n            for p in group['params']:\n                if p.grad is None:\n                    continue\n\n                state = self.state[p]\n                if len(state) == 0:\n                    state['exp_avg'] = torch.zeros_like(p)\n                    state['exp_avg_sq'] = torch.zeros_like(p)\n                    state['exp_avg_diff'] = torch.zeros_like(p)\n\n                grad = p.grad.mul_(clip_global_grad_norm)\n                if 'pre_grad' not in state or group['step'] == 1:\n                    state['pre_grad'] = grad\n\n                copy_grad = grad.clone()\n\n                exp_avg, exp_avg_sq, exp_avg_diff = state['exp_avg'], state['exp_avg_sq'], state['exp_avg_diff']\n                diff = grad - state['pre_grad']\n\n                update = grad + beta2 * diff\n                exp_avg.mul_(beta1).add_(grad, alpha=1 - beta1)  # m_t\n                exp_avg_diff.mul_(beta2).add_(diff, alpha=1 - beta2)  # diff_t\n                exp_avg_sq.mul_(beta3).addcmul_(update, update, value=1 - beta3)  # n_t\n\n                denom = ((exp_avg_sq).sqrt() / math.sqrt(bias_correction3)).add_(group['eps'])\n                update = ((exp_avg / bias_correction1 + beta2 * exp_avg_diff / bias_correction2)).div_(denom)\n\n                if group['no_prox']:\n                    p.data.mul_(1 - group['lr'] * group['weight_decay'])\n                    p.add_(update, alpha=-group['lr'])\n                else:\n                    p.add_(update, alpha=-group['lr'])\n                    p.data.div_(1 + group['lr'] * group['weight_decay'])\n\n                state['pre_grad'] = copy_grad","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:43:16.311430Z","iopub.execute_input":"2024-04-03T03:43:16.311727Z","iopub.status.idle":"2024-04-03T03:43:16.337183Z","shell.execute_reply.started":"2024-04-03T03:43:16.311702Z","shell.execute_reply":"2024-04-03T03:43:16.336366Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Helper functions\n# ====================================================\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\n\ndef asMinutes(s):\n    m = math.floor(s / 60)\n    s -= m * 60\n    return '%dm %ds' % (m, s)\n\n\ndef timeSince(since, percent):\n    now = time.time()\n    s = now - since\n    es = s / (percent)\n    rs = es - s\n    return '%s (remain %s)' % (asMinutes(s), asMinutes(rs))\n\n\ndef train_fn(fold, train_loader, teacher_model, model, criterion, optimizer, epoch, scheduler, device):\n    model.train()\n    scaler = torch.cuda.amp.GradScaler(enabled=CFG.apex)\n    losses = AverageMeter()\n    start = end = time.time()\n    global_step = 0\n    for step, batch in enumerate(train_loader):\n        spectrogram = batch['spectrogram'].to(device)\n        labels = batch['labels'].to(device)\n        batch_size = labels.size(0)\n        with torch.no_grad():\n            teacher_features = teacher_model.module.custom_layers[:-1](spectrogram)\n            teacher_preds = teacher_model(spectrogram)\n        with torch.cuda.amp.autocast(enabled=CFG.apex):\n            contrastive_target = torch.zeros(teacher_features.size(0)).to(device)\n            student_preds= model(spectrogram)\n            student_features = model.module.custom_layers[:-1](spectrogram)\n            loss = criterion(teacher_features, student_features, teacher_preds[:,:-1], student_preds[:,:-1], labels[:, :-1], contrastive_target)\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        losses.update(loss.item(), batch_size)\n        scaler.scale(loss).backward()\n        grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.max_grad_norm)\n        if (step + 1) % CFG.gradient_accumulation_steps == 0:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n            global_step += 1\n            if CFG.batch_scheduler:\n                scheduler.step()\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(train_loader)-1):\n            print('Epoch: [{0}][{1}/{2}] '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  'Grad: {grad_norm:.4f}  '\n                  'LR: {lr:.8f}  '\n                  .format(epoch+1, step, len(train_loader), \n                          remain=timeSince(start, float(step+1)/len(train_loader)),\n                          loss=losses,\n                          grad_norm=grad_norm,\n                          lr=scheduler.get_lr()[0]))\n        if CFG.wandb:\n            wandb.log({f\"[fold{fold}] loss\": losses.val,\n                       f\"[fold{fold}] lr\": scheduler.get_lr()[0]})\n    return losses.avg\n\n\ndef valid_fn(valid_loader, model, criterion, device):\n    losses = AverageMeter()\n    model.eval()\n    preds = []\n    start = end = time.time()\n    for step, batch in enumerate(valid_loader):\n        spectrogram = batch['spectrogram'].to(device)\n        labels = batch['labels'].to(device)\n        batch_size = labels.size(0)\n        with torch.no_grad():\n            y_preds = model(spectrogram)\n            loss = criterion(F.log_softmax(y_preds[:, 0:-1], dim=1), labels[:, 0:-1])\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        losses.update(loss.item(), batch_size)\n        preds.append(nn.Softmax(dim=1)(y_preds[:, 0:-1]).to('cpu').numpy())\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(valid_loader)-1):\n            print('EVAL: [{0}/{1}] '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  .format(step, len(valid_loader),\n                          loss=losses,\n                          remain=timeSince(start, float(step+1)/len(valid_loader))))\n    predictions = np.concatenate(preds)\n    return losses.avg, predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:43:16.338303Z","iopub.execute_input":"2024-04-03T03:43:16.338576Z","iopub.status.idle":"2024-04-03T03:43:16.363104Z","shell.execute_reply.started":"2024-04-03T03:43:16.338553Z","shell.execute_reply":"2024-04-03T03:43:16.362324Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train Loop","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# train loop\n# ====================================================\ndef train_loop(folds, fold, directory):\n    \n    LOGGER.info(f\"========== fold: {fold} training ==========\")\n\n    # ====================================================\n    # loader\n    # ====================================================\n    if CFG.stage1_pop1:\n        train_folds = folds[(folds['fold'] != fold)].reset_index(drop=True)\n    else:\n        train_folds = folds[(folds['fold'] != fold) & (folds['total_evaluators'] >= 10)].reset_index(drop=True)\n    valid_folds = folds[folds['fold'] == fold].reset_index(drop=True)\n    valid_labels = valid_folds[ CFG.target_cols].values\n    \n    train_dataset = CustomDataset(train_folds, augment=True, mode=\"train\")\n    valid_dataset = CustomDataset(valid_folds, augment=False, mode=\"train\")\n\n    train_loader = DataLoader(train_dataset,\n                              batch_size=CFG.batch_size,\n                              shuffle=True,\n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=True)\n    valid_loader = DataLoader(valid_dataset,\n                              batch_size=CFG.batch_size * 2,\n                              shuffle=False,\n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=False)\n\n    # ====================================================\n    # model & optimizer\n    # ====================================================\n    if CFG.stage2_pop2:\n        teacher_model = CustomModel(CFG)\n        \n        model_weight = \"/kaggle/input/stage1-effb0-zimmer-specs-train/pop_1_weight_oof/\" + f\"tf_effb0_ns_zimmer_specs_final_fold{fold}_best_version1_stage1.pth\"\n        checkpoint = torch.load(model_weight, map_location=device)\n        teacher_model.load_state_dict(checkpoint[\"model\"])\n        \n        for param in teacher_model.parameters():\n            param.requires_grad = False\n        teacher_model.eval()\n        teacher_model.to(device)\n        if CFG.t4_gpu:\n            teacher_model = nn.DataParallel(teacher_model)\n        \n    # CPMP: wrap the model to use all GPUs\n    student_model = CustomModel(CFG)\n    student_model.to(device)\n    if CFG.t4_gpu:\n        student_model = nn.DataParallel(student_model)\n    \n    def build_optimizer(cfg, model, device):\n        lr = cfg.lr\n        # lr = default_configs[\"lr\"]\n        if cfg.optimizer == \"SAM\":\n            base_optimizer = torch.optim.SGD  # define an optimizer for the \"sharpness-aware\" update\n            optimizer_model = SAM(model.parameters(), base_optimizer, lr=lr, momentum=0.9, weight_decay=cfg.weight_decay, adaptive=True)\n        elif cfg.optimizer == \"Ranger21\":\n            optimizer_model = Ranger21(model.parameters(), lr=lr, weight_decay=cfg.weight_decay, \n            num_epochs=cfg.epochs, num_batches_per_epoch=len(train_loader))\n        elif cfg.optimizer == \"SGD\":\n            optimizer_model = torch.optim.SGD(model.parameters(), lr=lr, weight_decay=cfg.weight_decay, momentum=0.9)\n        elif cfg.optimizer == \"Adam\":\n            optimizer_model = Adam(model.parameters(), lr=lr, weight_decay=CFG.weight_decay)\n        elif cfg.optimizer == \"AdamW\":\n            optimizer_model = AdamW(model.parameters(), lr=lr, weight_decay=CFG.weight_decay)\n        elif cfg.optimizer == \"Lion\":\n            optimizer_model = Lion(model.parameters(), lr=lr, weight_decay=cfg.weight_decay)\n        elif cfg.optimizer == \"Adan\":\n            optimizer_model = Adan(model.parameters(), lr=lr, weight_decay=cfg.weight_decay)\n    \n        return optimizer_model\n    \n    optimizer = build_optimizer(CFG, student_model, device)\n    \n    # ====================================================\n    # scheduler\n    # ====================================================\n    # ====================================================\n\n    def get_scheduler(optimizer):\n        if CFG.scheduler=='ReduceLROnPlateau':\n            scheduler = ReduceLROnPlateau(optimizer, **CFG.reduce_params)\n        elif CFG.scheduler=='CosineAnnealingLR':\n            scheduler = CosineAnnealingLR(optimizer, **CFG.cosanneal_params)\n        elif CFG.scheduler=='CosineAnnealingWarmRestarts':\n            scheduler = CosineAnnealingWarmRestarts(optimizer, **CFG.cosanneal_res_params)\n        elif CFG.scheduler=='OneCycleLR':\n            steps_per_epoch=len(train_loader),\n            scheduler = OneCycleLR(optimizer=optimizer, epochs=CFG.epochs, anneal_strategy=\"cos\", pct_start=0.05, steps_per_epoch=len(train_loader),\n        max_lr=CFG.lr, final_div_factor=100)\n        return scheduler\n    \n    scheduler = get_scheduler(optimizer)\n\n    # ====================================================\n    # loop\n    # ====================================================\n    train_loss = CustomLoss(weights=CFG.weights)\n    valid_loss = nn.KLDivLoss(reduction=\"batchmean\")\n    best_score = np.inf\n\n    for epoch in range(CFG.epochs):\n\n        start_time = time.time()\n\n        # train\n        avg_loss = train_fn(fold, train_loader, teacher_model, student_model, train_loss, optimizer, epoch, scheduler, device)\n\n        # eval\n        avg_val_loss, predictions = valid_fn(valid_loader, student_model, valid_loss, device)\n\n        elapsed = time.time() - start_time\n\n        LOGGER.info(f'Epoch {epoch+1} - avg_train_loss: {avg_loss:.4f}  avg_val_loss: {avg_val_loss:.4f}  time: {elapsed:.0f}s')\n        if CFG.wandb:\n            wandb.log({f\"[fold{fold}] epoch\": epoch+1, \n                       f\"[fold{fold}] avg_train_loss\": avg_loss, \n                       f\"[fold{fold}] avg_val_loss\": avg_val_loss,\n                       f\"[fold{fold}] score\": score})\n        \n        if best_score > avg_val_loss:\n            best_score = avg_val_loss\n            LOGGER.info(f'Epoch {epoch+1} - Save Best valid loss: {avg_val_loss:.4f} Model')\n            # CPMP: save the original model. It is stored as the module attribute of the DP model.\n            if CFG.stage1_pop1:\n                state_dict = student_model.module.state_dict() if CFG.t4_gpu else model.state_dict()\n                torch.save({'model': state_dict,\n                            'predictions': predictions},\n                             directory+f\"{CFG.model_name}_fold{fold}_best_version{VERSION}_stage1.pth\")\n            else:\n                state_dict = student_model.module.state_dict() if CFG.t4_gpu else model.state_dict()\n                torch.save({'model': state_dict,\n                            'predictions': predictions},\n                             directory+f\"{CFG.model_name}_fold{fold}_best_version{VERSION}_stage2.pth\")\n                \n    if CFG.stage1_pop1:\n        predictions = torch.load(directory+f\"{CFG.model_name}_fold{fold}_best_version{VERSION}_stage1.pth\", \n                             map_location=torch.device('cpu'))['predictions']\n    else:\n        predictions = torch.load(directory+f\"{CFG.model_name}_fold{fold}_best_version{VERSION}_stage2.pth\", \n                             map_location=torch.device('cpu'))['predictions']\n    valid_folds[[f\"pred_{c}\" for c in CFG.target_cols]] = predictions\n    valid_folds[CFG.target_cols] = valid_labels \n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    return valid_folds, best_score","metadata":{"execution":{"iopub.status.busy":"2024-04-03T03:43:16.364449Z","iopub.execute_input":"2024-04-03T03:43:16.364785Z","iopub.status.idle":"2024-04-03T03:43:16.393855Z","shell.execute_reply.started":"2024-04-03T03:43:16.364756Z","shell.execute_reply":"2024-04-03T03:43:16.393103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if __name__ == '__main__':\n    \n    if CFG.wandb:\n        \n        run = wandb.init(project='HMS_competition', \n                     name=CFG.model_name + f'version{VERSION}' + '_stage2',\n                     config=class2dict(CFG),\n                     group=CFG.model_name + '_stage2',\n                     job_type=\"train\",\n                     )\n    \n    if CFG.train:\n        oof_df = pd.DataFrame()\n        scores = []\n        for fold in range(CFG.n_fold):\n            if fold in CFG.trn_fold:\n                _oof_df, score = train_loop(train, fold, POP_2_DIR)\n                oof_df = pd.concat([oof_df, _oof_df])\n                scores.append(score)\n                LOGGER.info(f\"========== fold: {fold} result ==========\")\n                LOGGER.info(f'Score with best loss weights stage2: {score}')\n        oof_df = oof_df.reset_index(drop=True)\n        LOGGER.info(f\"========== CV ==========\")\n        LOGGER.info(f'Score with best loss weights stage2: {np.mean(scores)}')\n        oof_df.to_csv(POP_2_DIR+f'{CFG.model_name}_oof_df_version{VERSION}_stage2.csv', index=False)\n        \n    if CFG.wandb:\n        wandb.finish()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-04-03T03:43:16.394925Z","iopub.execute_input":"2024-04-03T03:43:16.395185Z"},"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\n\n# === Pre-process OOF ===\nlabel_cols = CFG.target_cols\ngt = oof_df[[\"eeg_id\"] + CFG.target_cols]\ngt.sort_values(by=\"eeg_id\", inplace=True)\ngt.reset_index(inplace=True, drop=True)\n\npreds = oof_df[[\"eeg_id\"] + CFG.pred_cols]\npreds.columns = [\"eeg_id\"] + CFG.target_cols\npreds.sort_values(by=\"eeg_id\", inplace=True)\npreds.reset_index(inplace=True, drop=True)\n\ny_trues = gt[CFG.target_cols]\ny_preds = preds[CFG.target_cols]\n\noof = pd.DataFrame(y_preds.copy())\noof['id'] = np.arange(len(oof))\n\ntrue = pd.DataFrame(y_trues.copy())\ntrue['id'] = np.arange(len(true))\n\ncv = score(solution=true, submission=oof, row_id_column_name='id')\nprint(f'CV Stage2 Score with {CFG.model_name} Spectrogram =',cv)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}