{"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":"none","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":160940372,"sourceType":"kernelVersion"}],"dockerImageVersionId":30635,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Harmful Brain Activity Classification - Pytorch Lightning Starter\n\nThis notebook provides a basic approach for classifying brain activity using a Deep Learning model based on 1-D Convolutions of EEGs. Use this as a starting point for experimenting with your models.\n\nThe main features of this notebook:\n\n- CONFIG-driven modeling.\n- PyTorch Lightning framework.\n- WANDB visualization.\n- Crossvalidation folds are used as an ensemble model.\n\n## Using this notebook\n\nCreate an account on wandb.ai and get an API key. Now attach the API key as a secret to this notebook (by clicking on Add-ons and Secrets) named `WANDB_API_KEY`. Also record your userid as key `WANDB_ENTITY`. This step is strictly optional, but it really helps to visualize the performance of your model. If you don't use this feature of the notebook then you can keep the Internet off.\n\nRun the whole notebook or save a version to do a full run and save a model. The model will be saved in the output directory. Now, you can turn off the Internet and submit the notebook. When you submit to the competition, it will reuse the saved model for making predictions on the test dataset. There is a separate subdirectory created for each experiment, so they will not overwrite each other. Feel free to delete any old experiment directories.\n\nIf you make model changes then first set the `development` option to true before running the notebook. This will do a quick training pass which will help identify any compilation errors. Once you have fixed your code then set development to false, assign a new experiment name, and run the notebook.\n\nAll of the model and training parameters are defined in the CONFIG string in the next section. Each parameter has been fully documented to help you understand what they do. Feel free to add more parameters as needed.\n\n## Credits\n- I copied a number of concepts from [this notebook](https://www.kaggle.com/code/tubotubo/cmi-submit) by [213tubo](https://www.kaggle.com/tubotubo).","metadata":{}},{"cell_type":"markdown","source":"# CONFIG","metadata":{}},{"cell_type":"code","source":"CONFIG_STR = \"\"\"\n# Set development to false when the code compiles and runs correctly.\ndevelopment: false\n\n# Give a unique name to each experiment\nexperiment: exp016\n\n# The notes are recorded in WANDB, so put something descriptive to\n# remind yourself what was changed in this version.\nnotes: split by label_id\n\ndata_dir: /kaggle/input/hms-harmful-brain-activity-classification\noutput_dir: /kaggle/working\ntemp_dir: /kaggle/temp\n\n# Setting processed_data_dir to a non-null value will make the trainer\n# first copy over all of the parquet files to numpy files in this directory.\n# This is a slow operation but makes each epoch run much faster.\nprocessed_data_dir: null # /kaggle/temp/processed_data\n\n# If you want to reuse a model from a previous run of this notebook\n# then \"Add Data\" from that previous version of this notebook.\n# This helps to speed up submissions.\nprev_notebook_ver: /kaggle/input/pytorch-lightning-starter-with-wandb-visualization\n\nseed: 42 # A random number generator seed for reproducibility.\n\nmodel:\n    class_name: BasicConvolution\n    \n    # the following parameters are passed to this class\n    BasicConvolution:\n        num_conv_layers: 16\n        num_kernels: 64\n        kernel_size: 3\n        stride: 1\n        dilation: 1\n        activation: \"tanh\" # relu, tanh \n        fcn_nodes: 128 # the number of nodes in a fully connected layer\n        dropout1_prob: 0\n        dropout2_prob: 0\n    \ntrain:\n    # Specify the column to split the dataset by. For example,\n    # label_id, eeg_id, patient_id, etc.\n    split_by_col: label_id\n    \n    # If one_per_split_by_col is true then in each epoch of\n    # training and validation we will not repeat the split_by_col.\n    # For example if we split by eeg_id then each training epoch\n    # will pick one random example for each eeg_id.\n    one_per_split_by_col: true\n\n    num_folds: 5  # The number of validation folds\n    \n    batch_size: 256  # The batch size used in training, validation and testing\n    num_workers: 4 # The number of workers used for loading the data.\n    accelerator: auto # Use GPU if available\n    precision: 32 # 16-mixed => automatic mixed precision (AMP). Or try `32`.\n    gradient_clip_val: 1.0\n    accumulate_grad_batches: 1\n    check_val_every_n_epoch: 1\n    deterministic: true\n    impute_zero: true # change missing values to zero o.w. change to mean\n    max_epochs: 17 # The number of training epochs\n    \n    # Restrict the maximum amount of time used for training.\n    # This only applies to one fold.\n    max_time: \"00:01:00:00\" # \"00:01:00:00\" => 1 hour, null => no limit\n    limit_val_batches: 1.0 # 1.0 => no limit, .25 => one fourth, 3 => exactly 3\n    limit_train_batches: 1.0 # 1.0 => no limit, .25 => one fourth, 3 => exactly 3\n\noptimizer:\n  lr: 0.0005  # The initial learning rate\n  \n  class_name: AdamW  # SGD or AdamW\n  SGD:\n    momentum: 0.9\n    weight_decay: 0\n    nesterov: false\n    dampening: 0\n  AdamW:\n    weight_decay: 0.01\n    \n  scheduler: cosine  # scheduler for decreasing the learning rate.\n  num_warmup_steps: 0\n\n    \n\"\"\"\nprint(CONFIG_STR)","metadata":{"execution":{"iopub.status.busy":"2024-01-29T23:57:55.762887Z","iopub.execute_input":"2024-01-29T23:57:55.763326Z","iopub.status.idle":"2024-01-29T23:57:55.771640Z","shell.execute_reply.started":"2024-01-29T23:57:55.763295Z","shell.execute_reply":"2024-01-29T23:57:55.770712Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions\n\n- configure_logger\n- info\n- json_to_py\n- wandb_login\n- impute_mean\n- impute_zero\n- create_directory\n- parquet_to_numpy\n- merge_csvs","metadata":{}},{"cell_type":"code","source":"from types import SimpleNamespace\nimport logging\nimport json\nimport yaml\nimport wandb\nimport os\nfrom kaggle_secrets import UserSecretsClient\nimport pandas as pd\nfrom pathlib import Path\nimport numpy as np\nimport joblib\nfrom tqdm.notebook import tqdm\nfrom typing import List\n\ndef configure_logger(level = logging.INFO):\n    root_logger = logging.getLogger()\n    # Kaggle notebooks have a builtin File logger, we ignore that\n    if len(root_logger.handlers) < 2:\n        # Configure logging with a custom formatter\n        formatter = logging.Formatter(\n            \"%(asctime)s - %(levelname)s - %(name)s - %(message)s\",\n            datefmt=\"%Y-%m-%d %H:%M:%S\",\n        )\n\n        # Create a handler and set the formatter\n        handler = logging.StreamHandler()\n        handler.setFormatter(formatter)\n\n        # Add the handler to the root logger\n        root_logger.addHandler(handler)\n\n        # Set the desired logging level\n        root_logger.setLevel(level)\n\ndef info(module: str, message: str):\n    logger = logging.getLogger(module)\n    logger.info(message)\n\ndef json_to_py(json_cfg):\n    \"\"\"\n    Convert a JSON object to a Python object, i.e.\n    insead of `j[\"foo\"][\"bar\"]` we can write `p.foo.bar`\n    where `p = json_to_py(j)`.\n    \"\"\"\n    return json.loads(json.dumps(json_cfg), object_hook=lambda d: SimpleNamespace(**d))\n\ndef wandb_login() -> bool:\n    \"\"\"Returns true if login was successful.\"\"\"\n    try:\n        user_secrets = UserSecretsClient()\n        os.environ[\"WANDB_API_KEY\"] = user_secrets.get_secret(\"WANDB_API_KEY\")\n        os.environ[\"WANDB_ENTITY\"] = user_secrets.get_secret(\"WANDB_ENTITY\")\n        #os.environ[\"WANDB_PROJECT\"] = \"HMS - Harmful Brain Activity Classification\"\n        return wandb.login()\n    except:\n        return False\n\ndef impute_mean(df: pd.DataFrame) -> pd.DataFrame:\n    return df.apply(lambda col: col.fillna(col.mean()), axis=0)\n\ndef impute_zero(df: pd.DataFrame) -> pd.DataFrame:\n    return df.apply(lambda col: col.fillna(0), axis=0)\n\ndef create_directory(directory_path) -> bool:\n    # Check if the directory already exists and returns False\n    if not os.path.exists(directory_path):\n        # Create the directory if it doesn't exist\n        os.makedirs(directory_path)\n        return True\n    else:\n        return False\n\ndef parquet_to_numpy(src_dir, dest_dir, impute_zero:bool):\n    \"\"\"\n    Copy all the parquet files in src_dir to dest_dir\n    inspired from this notebook:\n    https://www.kaggle.com/code/awsaf49/hms-hbac-kerascv-starter-notebook?scriptVersionId=160593469&cellId=18\n    \"\"\"\n    create_directory(dest_dir)\n    \n    def copy_one_file(src_file, dest_dir, impute_zero):\n        df = pd.read_parquet(src_file)\n        if impute_zero:\n            arr = impute_zero(df).to_numpy()\n        else:\n            arr = impute_mean(df).to_numpy()\n        prefix = src_file.name.split(\".\")[0]\n        np.save(Path(dest_dir) / (prefix + \".npy\"), arr)\n        \n    all_files = list(Path(src_dir).glob(\"*.parquet\"))\n    joblib.Parallel(n_jobs=-1, backend=\"loky\")(\n        joblib.delayed(copy_one_file)(filename, dest_dir, impute_zero)\n        for filename in tqdm(all_files)\n    )\n\n    import pandas as pd\n\n# Assuming submissions is a list of file names\nsubmissions = ['file1.csv', 'file2.csv', 'file3.csv']\n\ndef merge_csvs(src_csvs: List[str], tgt_csv: str, gby_col:str):\n    \"\"\"\n    Read all the `src_csvs`, concatenate them, group by the\n    `gby_col` and compute the mean of all other columns then\n    write out the result into `tgt_csv`.\n    \"\"\"\n    # Create an empty list to store individual DataFrames\n    dfs = []\n\n    # Iterate through each file in submissions and load the DataFrames\n    for file in src_csvs:\n        df = pd.read_csv(file)\n        dfs.append(df)\n\n    # Concatenate all DataFrames into a single DataFrame\n    result_df = pd.concat(dfs, ignore_index=True)\n\n    # Group by 'eeg_id' and calculate the mean of other columns\n    result_df = result_df.groupby('eeg_id').mean().reset_index()\n\n    # Write the resulting DataFrame to a CSV file\n    result_df.to_csv(tgt_csv, index=False)\n","metadata":{"execution":{"iopub.status.busy":"2024-01-29T23:57:55.773808Z","iopub.execute_input":"2024-01-29T23:57:55.774186Z","iopub.status.idle":"2024-01-29T23:57:55.793967Z","shell.execute_reply.started":"2024-01-29T23:57:55.774153Z","shell.execute_reply":"2024-01-29T23:57:55.793167Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Module\n\n- Datasets\n    - OffsetEEG\n    - FullEEG\n- DataLoader\n    - BACDataModule","metadata":{}},{"cell_type":"code","source":"import pytorch_lightning as ptl\nfrom pathlib import Path, PosixPath\nimport logging\nimport pandas as pd\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import DataLoader, Dataset\nimport pyarrow.parquet as pq\nimport numpy as np\nfrom typing import Optional\n\nEEG_SAMPRATE = 200 # samples per second\nEEG_LENGTH = 50  # seconds\nEEG_FEATURES = 20 # number of columns of EEG data\nEEG_CLASSES = [\"seizure\", \"lpd\", \"gpd\", \"lrda\", \"grda\", \"other\"]\n\nclass OffsetEEG(Dataset):\n    \"\"\"\n    An offset within an EEG is returned along with the true label.\n    This dataset is used for training and validation.\n    \"\"\"\n    def __init__(self, df: pd.DataFrame, eeg_dir: PosixPath, spec_dir: PosixPath,\n                 processed_data: bool, impute_zero:bool, one_per_col: Optional[str]):\n        super().__init__()\n        self.df = df\n        self.eeg_dir = eeg_dir\n        self.spec_dir = spec_dir\n        self.processed_data = processed_data\n        self.impute_zero = impute_zero\n        self.one_per_col = one_per_col\n        if self.one_per_col is not None:\n            self.unique_ids = self.df[self.one_per_col].unique()\n        # expected number of rows in each data point\n        self.num_rows = EEG_SAMPRATE * EEG_LENGTH\n    \n    def __len__(self):\n        if self.one_per_col is not None:\n            return len(self.unique_ids)\n        else:\n            return len(self.df)\n    \n    def __getitem__(self, index):\n        if self.one_per_col is not None:\n            # return a random row for the given patient\n            col_id = self.unique_ids[index]\n            row = self.df[self.df[self.one_per_col] == col_id].sample().iloc[0]\n        else:\n            row = self.df.iloc[index]\n        # extract the EEG at the given offset\n        offset = int(row.eeg_label_offset_seconds * EEG_SAMPRATE)\n        if self.processed_data:\n            eeg_file_path = self.eeg_dir / f\"{row.eeg_id}.npy\"\n            eeg = np.load(eeg_file_path)[offset: offset + self.num_rows]\n        else:\n            eeg_file_path = self.eeg_dir / f\"{row.eeg_id}.parquet\"\n            eeg_file = pq.ParquetFile(eeg_file_path)\n            df = eeg_file.read().slice(offset, self.num_rows).to_pandas()\n            if self.impute_zero:\n                eeg = impute_zero(df).to_numpy()\n            else:\n                eeg = impute_mean(df).to_numpy()\n        assert(eeg.shape == (self.num_rows, EEG_FEATURES))\n        # construct the true class label\n        label = np.array([getattr(row, f\"{cls}_vote\") for cls in EEG_CLASSES])\n        label = label / label.sum()\n        # EEG should have input features first as expected by torch\n        return dict(eeg=eeg.T, name=f\"{row.eeg_id}_{row.eeg_sub_id}\",\n                   target=label)\n\nclass FullEEG(Dataset):\n    \"\"\"\n    A full EEG file is returned.\n    \"\"\"\n    def __init__(self, df: pd.DataFrame, eeg_dir: PosixPath, spec_dir: PosixPath, impute_zero:bool):\n        super().__init__()\n        self.df = df\n        self.eeg_dir = eeg_dir\n        self.spec_dir = spec_dir\n        self.impute_zero = impute_zero\n        # expected number of rows in each data point\n        self.num_rows = EEG_SAMPRATE * EEG_LENGTH\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n        eeg_file_path = self.eeg_dir / f\"{row.eeg_id}.parquet\"\n        eeg_file = pq.ParquetFile(eeg_file_path)\n        assert(eeg_file.metadata.num_rows == EEG_SAMPRATE * EEG_LENGTH)\n        if self.impute_zero:\n            eeg = impute_zero(eeg_file.read().to_pandas()).to_numpy()\n        else:\n            eeg = impute_mean(eeg_file.read().to_pandas()).to_numpy()\n        assert(eeg.shape == (self.num_rows, EEG_FEATURES))\n        # EEG should have input features first as expected by torch\n        return dict(eeg=eeg.T, name=str(row.eeg_id))\n\nclass BACDataModule(ptl.LightningDataModule):\n\n    def __init__(self, cfg: SimpleNamespace):\n        super().__init__()\n        self.cfg = cfg\n        self.data_dir = Path(self.cfg.data_dir)\n        if self.cfg.processed_data_dir is not None:\n            self.train_eegs = Path(self.cfg.processed_data_dir) / \"train_eegs\"\n            self.train_specs = Path(self.cfg.processed_data_dir) / \"train_spectrograms\"\n            self.processed_data = True\n        else:\n            self.train_eegs = Path(self.data_dir) / \"train_eegs\"\n            self.train_specs = Path(self.data_dir) / \"train_spectrograms\"\n            self.processed_data = False    \n    \n    def prepare_data(self):\n        if self.processed_data and not self.train_eegs.exists():\n            info(\"BACDataModule\", f\"Processing training eegs into directory {self.train_eegs}\")\n            parquet_to_numpy(Path(self.data_dir) / \"train_eegs\", self.train_eegs,\n                             self.cfg.train.impute_zero)\n        # TODO: create self.train_specs    \n            \n    def setup(self, stage: str):\n        info(\"BACDataModule\", f\"setup {stage=}\")\n        if stage == \"fit\":\n            self.full_train_df = pd.read_csv(self.data_dir / \"train.csv\")\n            \n            # Get unique IDS for the split_by_col\n            unique_ids = self.full_train_df[self.cfg.train.split_by_col].unique()\n\n            validation_ids = unique_ids[self.cfg.train.fold_idx::self.cfg.train.num_folds]\n            validation_mask = self.full_train_df[self.cfg.train.split_by_col].isin(validation_ids)\n            self.train_df = self.full_train_df[~validation_mask]\n            self.val_df = self.full_train_df[validation_mask]\n\n            info(\n                \"BACDataModule\",\n                f\"Training fold {self.cfg.train.fold_idx} of {self.cfg.train.num_folds}\"\n                f\" folds has {len(self.train_df)} train labels\"\n                f\" and {len(self.val_df)} validation labels (split by {self.cfg.train.split_by_col}).\"\n            )\n        elif stage == \"test\":\n            self.test_df = pd.read_csv(self.data_dir / \"test.csv\")\n            info(\"BACDataModule\", f\"Testing on {len(self.test_df)} EEGs.\")\n        \n        else:\n            raise ValueError(f\"Unsupported {stage=}\")\n\n    def train_dataloader(self) -> DataLoader:\n        return DataLoader(\n            OffsetEEG(self.train_df, self.train_eegs, self.train_specs,\n                      self.processed_data, self.cfg.train.impute_zero,\n                      self.cfg.train.split_by_col if \n                      self.cfg.train.one_per_split_by_col else None),\n            batch_size=self.cfg.train.batch_size,\n            shuffle=True,\n            num_workers=self.cfg.train.num_workers,\n            pin_memory=True,\n            drop_last=True,\n        )\n    \n    def val_dataloader(self) -> DataLoader:\n        return DataLoader(\n            OffsetEEG(self.val_df, self.train_eegs, self.train_specs,\n                      self.processed_data, self.cfg.train.impute_zero,\n                      self.cfg.train.split_by_col if \n                      self.cfg.train.one_per_split_by_col else None),\n            batch_size=self.cfg.train.batch_size,\n            shuffle=False,\n            num_workers=self.cfg.train.num_workers,\n            pin_memory=True,\n            drop_last=False,\n        )\n    \n    def test_dataloader(self) -> DataLoader:\n        return DataLoader(\n            FullEEG(self.test_df,\n                    self.data_dir / \"test_eegs\",\n                    self.data_dir / \"test_spectrograms\", self.cfg.train.impute_zero),\n            batch_size=self.cfg.train.batch_size,\n            shuffle=False,\n            num_workers=self.cfg.train.num_workers,\n            pin_memory=True,\n            drop_last=False,\n        )\n","metadata":{"execution":{"iopub.status.busy":"2024-01-29T23:57:55.899831Z","iopub.execute_input":"2024-01-29T23:57:55.900163Z","iopub.status.idle":"2024-01-29T23:57:56.204237Z","shell.execute_reply.started":"2024-01-29T23:57:55.900138Z","shell.execute_reply":"2024-01-29T23:57:56.203229Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model Module\n\n- Models\n    - BasicConvolution\n- LightningModule\n    - BACModelModule","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport pytorch_lightning as ptl\nfrom transformers import get_cosine_schedule_with_warmup\nfrom types import SimpleNamespace\nfrom pathlib import Path\nimport pandas as pd\nimport wandb\nimport importlib\nfrom collections import OrderedDict\n\nclass BasicConvolution(nn.Module):\n    def __init__(self,\n                 input_len,\n                 input_features,\n                 output_classes,\n                 num_conv_layers,\n                 num_kernels,\n                 kernel_size,\n                 stride,\n                 dilation,\n                 activation,\n                 dropout1_prob,\n                 fcn_nodes,\n                 dropout2_prob,\n                ):\n        super().__init__()\n        layers = OrderedDict()\n        for l in range(num_conv_layers):\n            layers[f\"conv{l+1}\"] = nn.Conv1d(input_features, num_kernels, kernel_size=kernel_size,\n                                             stride=stride, dilation=dilation, bias=True)\n            layers[f\"bn{l+1}\"] = nn.BatchNorm1d(num_kernels)\n            if activation == \"relu\":\n                layers[f\"relu{l+1}\"] = nn.ReLU(inplace=True)\n            elif activation == \"tanh\":\n                layers[f\"tanh{l+1}\"] = nn.Tanh()\n            else:\n                raise ValueError(f\"Unknown {activation=}\")\n            input_features = num_kernels\n            # compute the input_len for the next layer\n            input_len = self.calculate_conv_output_size(input_len, kernel_size, stride, dilation)\n            # after every 2 layers add a max pooling layer\n            if (l+1) % 2 == 0:\n                layers[f\"max{l+1}\"] = nn.MaxPool1d(2)\n                input_len = self.calculate_conv_output_size(input_len, 2, 2, 1)\n        self.conv = nn.Sequential(layers)\n        self.dropout1 =  nn.Dropout1d(dropout1_prob) # drop an entire channel\n        self.fc1 = nn.Linear(num_kernels * input_len, fcn_nodes)\n        self.dropout2 = nn.Dropout(dropout2_prob)\n        self.fc2 = nn.Linear(fcn_nodes, output_classes)\n\n    def calculate_conv_output_size(self, input_length, kernel_size, stride, dilation):\n        # Assuming padding=0\n        return (input_length - dilation * (kernel_size - 1) - 1) // stride + 1\n    \n    def forward(self, x):\n        \"\"\"\n        x: N, F, T -> N, T, C\n\n        where\n          N - Batch size.\n          F - Number of input features.\n          T - Input length.\n          C - Number of output classes.\n        \"\"\"\n        x = self.conv(x)\n        x = self.dropout1(x)\n        # Flatten the output before passing through the linear layer\n        x = x.view(x.size(0), -1)\n\n        x = self.fc1(x)\n        x = F.relu(x)\n        x = self.dropout2(x)\n        x = self.fc2(x)\n        return x\n\nclass BACModelModule(ptl.LightningModule):\n    def __init__(self, cfg: SimpleNamespace):\n        super().__init__()\n        self.save_hyperparameters()\n        self.cfg = cfg\n        self.output_dir = Path(self.cfg.output_dir)\n        model_class_name = self.cfg.model.class_name\n        model_class = globals()[model_class_name]\n        self.model_params = vars(getattr(cfg.model, model_class_name))\n        # All models have the following parameters that are not configurable\n        self.model_params.update(dict(\n            input_len=EEG_SAMPRATE * EEG_LENGTH,\n            input_features=EEG_FEATURES,\n            output_classes=len(EEG_CLASSES)))\n        self.model = model_class(**self.model_params)\n        self.loss_fn = nn.KLDivLoss(reduction=\"batchmean\")\n\n    def forward(self, batch):\n        out = {}\n        logits = self.model(batch[\"eeg\"])\n        if torch.isnan(logits).any():\n            torch.save(logits, self.output_dir / \"logits.pt\")\n            torch.save(batch[\"eeg\"], self.output_dir / \"eeg.pt\")\n            torch.save(self.model.state_dict(), self.output_dir / \"model.ckpt\")\n            torch.save(self.model_params, self.output_dir / \"model.params\")\n            raise ValueError(\"logits are nan\")\n        logprobs = F.log_softmax(logits, dim=1)\n        if \"target\" in batch:\n            out[\"loss\"] = self.loss_fn(logprobs, batch['target'])\n            #print(f\"loss {out['loss'].item()}\")\n        out['prob'] = np.exp(logprobs.detach().cpu().numpy())\n        return out\n    \n    def on_train_epoch_start(self):\n        self.train_epoch_loss = 0.\n        self.train_epoch_cnt = 0\n        self.epoch_metrics = {}\n        self.epoch_conf_mat = None\n\n    def training_step(self, batch, batch_idx):\n        out = self.forward(batch)\n        self.train_epoch_loss += out[\"loss\"].item() * len(batch['name'])\n        self.train_epoch_cnt += len(batch['name'])\n        self.log(\"loss\", out['loss'].item(), batch_size=len(batch['name']),\n                 on_step=True, on_epoch=False, prog_bar=True, logger=True)\n        return out[\"loss\"]\n\n    def on_validation_epoch_start(self):\n        self.val_epoch_loss = 0.\n        self.val_epoch_cnt = 0\n        self.val_epoch_true = []\n        self.val_epoch_prob = []\n        self.val_epoch_name = []\n\n    def validation_step(self, batch, batch_idx):\n        out = self.forward(batch)\n        self.val_epoch_loss += out[\"loss\"].item() * len(batch['name'])\n        self.val_epoch_cnt += len(batch['name'])\n        self.val_epoch_true.extend(batch['target'].cpu().numpy())\n        self.val_epoch_prob.extend(out['prob'])\n        self.val_epoch_name.extend(batch[\"name\"])\n        return out[\"loss\"]\n\n    def on_validation_epoch_end(self):\n        self.epoch_metrics[\"val_loss\"] = self.val_epoch_loss / self.val_epoch_cnt\n        \n        # Create a data frame from all the validation examples\n        # and write it to disk.\n        df = pd.DataFrame(data=np.concatenate([np.reshape(self.val_epoch_name, (-1, 1)),\n                                               self.val_epoch_true,\n                                               self.val_epoch_prob], axis=1),\n                          columns=[\"name\"] + [f\"true_{cls}\" for cls in EEG_CLASSES]\n                          + [f\"prob_{cls}\" for cls in EEG_CLASSES])\n        df.to_csv(self.output_dir / \"val.csv\", index=False)\n        \n        self.epoch_conf_mat = wandb.plot.confusion_matrix(probs=np.array(self.val_epoch_prob),\n                                                          y_true=np.argmax(self.val_epoch_true, axis=1),\n                                                          class_names=EEG_CLASSES)\n        del self.val_epoch_loss\n        del self.val_epoch_cnt\n        del self.val_epoch_true\n        del self.val_epoch_prob\n        del self.val_epoch_name\n    \n    def on_train_epoch_end(self):\n        self.epoch_metrics[\"train_loss\"] = self.train_epoch_loss / self.train_epoch_cnt\n        \n        # First log the metrics and then the rest.\n        self.log_dict(\n            self.epoch_metrics, on_step=False, on_epoch=True, logger=True, prog_bar=True\n        )\n        if self.epoch_conf_mat is not None and self.cfg.wandb_enabled:\n            self.logger.experiment.log({\"conf_mat\": self.epoch_conf_mat})\n        \n        del self.train_epoch_loss\n        del self.train_epoch_cnt\n        del self.epoch_metrics\n        del self.epoch_conf_mat\n\n    def configure_optimizers(self):\n        optim_class_name = self.cfg.optimizer.class_name\n        optim_module = importlib.import_module(\".optim\", \"torch\")\n        optim_class = getattr(optim_module, optim_class_name)\n        extra_params = getattr(self.cfg.optimizer, optim_class_name, {})\n        optimizer = optim_class(\n            self.parameters(), lr=self.cfg.optimizer.lr, **vars(extra_params)\n        )\n        if not hasattr(self.cfg.optimizer, \"scheduler\"):\n            return optimizer\n        elif self.cfg.optimizer.scheduler.lower() != \"cosine\":\n            raise ValueError(\n                \"Unsupported scheduler '{}'\".format(self.cfg.optimizer.scheduler)\n            )\n        scheduler = get_cosine_schedule_with_warmup(\n            optimizer,\n            num_training_steps=self.trainer.estimated_stepping_batches,\n            num_warmup_steps=self.cfg.optimizer.num_warmup_steps,\n        )\n        return [optimizer], [\n            {\n                \"scheduler\": scheduler,\n                \"interval\": \"step\",\n                \"frequency\": 1,\n                \"name\": \"lr\",\n            },\n        ]\n    \n    def on_test_epoch_start(self):\n        self.test_names = []\n        self.test_probs = []\n\n    def test_step(self, batch, batch_idx):\n        out = self.forward(batch)\n        self.test_names.extend(batch['name'])\n        self.test_probs.extend(out['prob'])\n    \n    def on_test_epoch_end(self):\n        \"\"\"we will write out a submission.csv to the output dir\"\"\"\n        df = pd.DataFrame(data=np.concatenate([np.reshape(self.test_names, (-1, 1)),\n                                               self.test_probs], axis=1),\n                          columns=[\"eeg_id\"] + [f\"{cls}_vote\" for cls in EEG_CLASSES])\n        df.to_csv(self.output_dir / \"submission.csv\", index=False)\n        del self.test_names\n        del self.test_probs\n","metadata":{"execution":{"iopub.status.busy":"2024-01-29T23:57:56.206442Z","iopub.execute_input":"2024-01-29T23:57:56.206845Z","iopub.status.idle":"2024-01-29T23:57:56.245436Z","shell.execute_reply.started":"2024-01-29T23:57:56.206808Z","shell.execute_reply":"2024-01-29T23:57:56.244462Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Trainer\n- train\n- train_one_fold","metadata":{}},{"cell_type":"code","source":"import yaml\nimport wandb\nfrom pytorch_lightning import Trainer, seed_everything\nfrom pytorch_lightning.callbacks import LearningRateMonitor, RichModelSummary\nfrom pytorch_lightning.loggers import WandbLogger\nfrom kaggle_secrets import UserSecretsClient\nimport shutil\n\ndef train(config_str):\n    configure_logger()\n    json_cfg = yaml.safe_load(CONFIG_STR)\n    \n    # copy the output of the previous version of this notebook for continuity\n    if json_cfg[\"prev_notebook_ver\"] is not None:\n        shutil.copytree(json_cfg[\"prev_notebook_ver\"], json_cfg[\"output_dir\"], dirs_exist_ok=True)\n        \n    # configure wandb\n    json_cfg[\"wandb_enabled\"] = not json_cfg[\"development\"] and wandb_login()\n    if not json_cfg[\"wandb_enabled\"]:\n        info(\"train\", \"WANDB was disabled.\")\n\n    # In development mode we do a minimal 2-fold cross validation\n    if json_cfg[\"development\"]:\n        json_cfg[\"train\"][\"num_folds\"] = 2\n        #json_cfg[\"processed_data_dir\"] = None\n    \n    submissions = []\n    num_folds = json_cfg[\"train\"][\"num_folds\"]\n    for fold_idx in range(num_folds):\n        json_cfg[\"train\"][\"fold_idx\"] = fold_idx\n        info(\"train\", f\"Training fold {fold_idx} of {num_folds}\")\n        submissions.append(train_one_fold(json_cfg))\n    \n    merge_csvs(submissions, Path(json_cfg[\"output_dir\"]) / \"submission.csv\", \"eeg_id\")\n\ndef train_one_fold(json_cfg):\n    cfg = json_to_py(json_cfg)\n    \n    # Each fold gets its own experiment name and directories\n    exp_name = f\"{cfg.experiment}_{cfg.train.fold_idx}\"\n    cfg.output_dir = Path(cfg.output_dir) / exp_name\n    cfg.temp_dir = Path(cfg.temp_dir) /exp_name\n    # and a unique seed\n    cfg.seed = cfg.seed + 13 * cfg.train.fold_idx\n    \n    seed_everything(cfg.seed)\n    create_directory(cfg.output_dir)\n    create_directory(cfg.temp_dir)\n\n    lr_monitor = LearningRateMonitor(\"epoch\")\n    model_summary = RichModelSummary(max_depth=3)\n\n    # If this experiment was already run then we don't have to re-run this on wandb.\n    ckpt_path = Path(cfg.output_dir) / \"trainer.ckpt\"\n    if ckpt_path.exists():\n        cfg.wandb_enabled = False\n\n    run = wandb.init(\n        job_type=f\"fold {cfg.train.fold_idx} of {cfg.train.num_folds}\",\n        dir=cfg.temp_dir,\n        config=json_cfg,\n        project=\"HMS - Harmful Brain Activity Classification\",\n        reinit=True,\n        group=cfg.experiment,\n        name=exp_name,\n        notes=cfg.notes,\n        save_code=cfg.train.fold_idx==0, # save the code only in the first fold\n        mode=\"disabled\" if not cfg.wandb_enabled else \"online\",\n    )\n    pl_logger = WandbLogger(experiment=run)\n\n    data = BACDataModule(cfg)\n    model = BACModelModule(cfg)\n\n    trainer = Trainer(\n        default_root_dir=cfg.temp_dir,\n        deterministic=cfg.train.deterministic,\n        # num_nodes=cfg.train.num_gpus,\n        accelerator=cfg.train.accelerator,\n        precision=cfg.train.precision,\n        gradient_clip_val=cfg.train.gradient_clip_val,\n        accumulate_grad_batches=cfg.train.accumulate_grad_batches,\n        # Note that RichProgressBar doesn't work well on kaggle hence using the\n        # builtin tqdm progress bar.\n        callbacks=[lr_monitor, model_summary],\n        logger=pl_logger,\n        num_sanity_val_steps=0,\n        sync_batchnorm=True,\n        check_val_every_n_epoch=1,\n        max_time=cfg.train.max_time,\n        # During development we only run for 2 epochs with 2 train and val batches.\n        log_every_n_steps=1 if cfg.development else 5,\n        max_epochs=2 if cfg.development else cfg.train.max_epochs,\n        limit_train_batches=2 if cfg.development else cfg.train.limit_train_batches,\n        limit_val_batches=2 if cfg.development else cfg.train.limit_val_batches,\n    )\n\n    # Load from a previous notebook if possible.\n    if ckpt_path.exists():\n        info(\"train\", f\"Reusing existing checkpoint {ckpt_path}\")\n        trainer.fit(model, data, ckpt_path=ckpt_path)\n    else:\n        trainer.fit(model, data)\n        info(\"train\", f\"Checkpointing model to {ckpt_path}.\")\n        trainer.save_checkpoint(ckpt_path)\n    \n    # Write the submission.csv file.\n    trainer.test(model, datamodule=data)\n\n    run.finish()\n\n    return cfg.output_dir / \"submission.csv\"","metadata":{"execution":{"iopub.status.busy":"2024-01-29T23:57:56.247030Z","iopub.execute_input":"2024-01-29T23:57:56.247401Z","iopub.status.idle":"2024-01-29T23:57:56.266552Z","shell.execute_reply.started":"2024-01-29T23:57:56.247344Z","shell.execute_reply":"2024-01-29T23:57:56.265639Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Finally, run the training","metadata":{}},{"cell_type":"code","source":"train(CONFIG_STR)","metadata":{"execution":{"iopub.status.busy":"2024-01-29T23:57:56.269328Z","iopub.execute_input":"2024-01-29T23:57:56.269938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# [Optional] Debugging Code","metadata":{}},{"cell_type":"code","source":"if False:\n    import torch\n    eeg = torch.load(\"eeg.pt\")\n    eeg.shape, torch.isnan(eeg).any()\n    params = torch.load(\"model.params\")\n    model = BasicConvolution(**params)\n    model.load_state_dict(torch.load(\"model.ckpt\"))\n    model.cuda()\n    model.eval()\n    torch.isnan(model(eeg)).any()","metadata":{"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if False:\n    # load the config options\n    JSON_CONFIG = yaml.safe_load(CONFIG_STR)\n    CONFIG = json_to_py(JSON_CONFIG)\n    CONFIG.development = False\n    CONFIG.num_workers = 20\n    CONFIG.train.batch_size = 256\n    data = BACDataModule(CONFIG)\n    data.setup(\"fit\")\n    \n    #for bat in tqdm.tqdm_notebook(data.train_dataloader()):\n    #    pass","metadata":{"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if False:\n    bat = next(iter(data.val_dataloader()))\n    bat.keys(), bat['eeg'].shape, bat['target'].shape, len(bat['name']), bat['name'][0]","metadata":{"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if False:\n    model = BasicConvolution(input_len=EEG_SAMPRATE * EEG_LENGTH, input_features=EEG_FEATURES, output_classes=len(EEG_CLASSES), num_kernels=10, kernel_size=3)\n    out = model(bat['eeg'])\n    out.shape, bat['target'].shape","metadata":{"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if False:\n    kl_loss = nn.KLDivLoss(reduction=\"batchmean\")\n    kl_loss(F.log_softmax(out, dim=1), bat['target']).item()","metadata":{"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if False:\n    bat['name'][15], bat['target'][15]","metadata":{"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if False:\n    data.val_df.iloc[15]","metadata":{"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]}]}