{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.13"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7392733,"sourceType":"datasetVersion","datasetId":4297749},{"sourceId":7392775,"sourceType":"datasetVersion","datasetId":4297782},{"sourceId":7447509,"sourceType":"datasetVersion","datasetId":4334995},{"sourceId":7585255,"sourceType":"datasetVersion","datasetId":4415285},{"sourceId":7637585,"sourceType":"datasetVersion","datasetId":4430326}],"dockerImageVersionId":30646,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os, gc\nos.environ[\"CUDA_VISIBLE_DEVICES\"]=\"0,1\"\nimport pandas as pd, numpy as np\nimport matplotlib.pyplot as plt\nimport torch\nimport numpy\nimport torchvision.models as models\n\nfrom brain_solver import preprocess_eeg_data\n\nVER = 3","metadata":{"execution":{"iopub.status.busy":"2024-02-16T15:02:46.924634Z","iopub.execute_input":"2024-02-16T15:02:46.925456Z","iopub.status.idle":"2024-02-16T15:02:56.687649Z","shell.execute_reply.started":"2024-02-16T15:02:46.925408Z","shell.execute_reply":"2024-02-16T15:02:56.685391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('data/hms-harmful-brain-activity-classification/train.csv')\n\nTARGETS = df.columns[-6:]\nprint('Train shape:', df.shape )\nprint('Targets', list(TARGETS))\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-02-16T15:02:56.688706Z","iopub.status.idle":"2024-02-16T15:02:56.689118Z","shell.execute_reply.started":"2024-02-16T15:02:56.688923Z","shell.execute_reply":"2024-02-16T15:02:56.68894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"##  Create Non-Overlapping Eeg Id Train Data\nThe competition data description says that test data does not have multiple crops from the same eeg_id. Therefore we will train and validate using only 1 crop per eeg_id. There is a discussion about this here.\nhttps://www.kaggle.com/competitions/hms-harmful-brain-activity-classification/discussion/467021","metadata":{}},{"cell_type":"code","source":"preprocess_eeg_data(df, TARGETS, consensus_col='expert_consensus')","metadata":{"execution":{"iopub.status.busy":"2024-02-16T15:02:56.691568Z","iopub.status.idle":"2024-02-16T15:02:56.692544Z","shell.execute_reply.started":"2024-02-16T15:02:56.692247Z","shell.execute_reply":"2024-02-16T15:02:56.692274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Create ResNet18 model","metadata":{}},{"cell_type":"code","source":"# Model creation\nmodel = torch.hub.load('pytorch/vision:v0.10.0', 'resnet18', pretrained=True) # Can change to 152 or some other version","metadata":{"execution":{"iopub.status.busy":"2024-02-16T15:02:56.694287Z","iopub.status.idle":"2024-02-16T15:02:56.695257Z","shell.execute_reply.started":"2024-02-16T15:02:56.694938Z","shell.execute_reply":"2024-02-16T15:02:56.694962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(net, data_loaders, epochs=100, lr=0.01, device=d2l.try_gpu()):\n    # Trains the model net with data from the data_loaders['train'] and data_loaders['val'].\n    net = net.to(device)\n    \n    optimizer = torch.optim.Adam(net.parameters(), lr=lr)\n\n    animator = d2l.Animator(xlabel='epoch',\n                            legend=['train loss', 'train acc', 'validation loss', 'validation acc'],\n                            figsize=(10, 5))\n\n    timer = {'train': d2l.Timer(), 'val': d2l.Timer()}\n\n    for epoch in range(epochs):\n        # monitor loss, accuracy, number of samples\n        metrics = {'train': d2l.Accumulator(3), 'val': d2l.Accumulator(3)}\n\n        for phase in ('train', 'val'):\n            # switch network to train/eval mode\n            net.train(phase == 'train')\n\n            for i, (x, y) in enumerate(data_loaders[phase]):\n                timer[phase].start()\n\n                # move to device\n                x = x.to(device)\n                y = y.to(device)\n\n                # compute prediction\n                y_hat = net(x)\n                \n                if y_hat.shape[1] == 1:\n                    # compute binary cross-entropy loss\n                    loss = torch.nn.BCEWithLogitsLoss()(y_hat[:, 0], y.to(torch.float))\n                else:\n                    # compute cross-entropy loss\n                    loss = torch.nn.CrossEntropyLoss()(y_hat, y)\n\n                if phase == 'train':\n                    # compute gradients and update weights\n                    optimizer.zero_grad()\n                    loss.backward()\n                    optimizer.step()\n\n                metrics[phase].add(loss * x.shape[0],\n                                   accuracy(y_hat, y) * x.shape[0],\n                                   x.shape[0])\n\n                timer[phase].stop()\n\n        animator.add(epoch + 1,\n            (metrics['train'][0] / metrics['train'][2],\n             metrics['train'][1] / metrics['train'][2],\n             metrics['val'][0] / metrics['val'][2],\n             metrics['val'][1] / metrics['val'][2]))\n\n    train_loss = metrics['train'][0] / metrics['train'][2]\n    train_acc  = metrics['train'][1] / metrics['train'][2]\n    val_loss   = metrics['val'][0] / metrics['val'][2]\n    val_acc    = metrics['val'][1] / metrics['val'][2]\n    examples_per_sec = metrics['train'][2] * epochs / timer['train'].sum()\n    \n    print(f'train loss {train_loss:.3f}, train acc {train_acc:.3f}, '\n          f'val loss {val_loss:.3f}, val acc {val_acc:.3f}')\n    print(f'{examples_per_sec:.1f} examples/sec '\n          f'on {str(device)}')","metadata":{"execution":{"iopub.status.busy":"2024-02-16T15:02:56.696857Z","iopub.status.idle":"2024-02-16T15:02:56.697287Z","shell.execute_reply.started":"2024-02-16T15:02:56.697086Z","shell.execute_reply":"2024-02-16T15:02:56.697102Z"},"trusted":true},"execution_count":null,"outputs":[]}]}