{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7447509,"sourceType":"datasetVersion","datasetId":4334995}],"dockerImageVersionId":30746,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n#for dirname, _, filenames in os.walk('/kaggle/input'):\n#    for filename in filenames:\n#        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-07-23T12:19:00.613167Z","iopub.execute_input":"2024-07-23T12:19:00.614111Z","iopub.status.idle":"2024-07-23T12:19:01.783668Z","shell.execute_reply.started":"2024-07-23T12:19:00.614075Z","shell.execute_reply":"2024-07-23T12:19:01.782611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\nimport math\nfrom tqdm.notebook import tqdm, trange\nimport time\n\nimport torch\nimport torchvision\nfrom torchvision.transforms import ToTensor\nfrom torch.utils.data import TensorDataset, DataLoader\nfrom torchvision.transforms import v2\nfrom sklearn.model_selection import train_test_split","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:19:01.785488Z","iopub.execute_input":"2024-07-23T12:19:01.786016Z","iopub.status.idle":"2024-07-23T12:19:08.425258Z","shell.execute_reply.started":"2024-07-23T12:19:01.785984Z","shell.execute_reply":"2024-07-23T12:19:08.424329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.environ['CUDA_LAUNCH_BLOCKING']=\"1\"\nos.environ['TORCH_USE_CUDA_DSA'] = \"1\"","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:19:32.132293Z","iopub.execute_input":"2024-07-23T12:19:32.133064Z","iopub.status.idle":"2024-07-23T12:19:32.137924Z","shell.execute_reply.started":"2024-07-23T12:19:32.133029Z","shell.execute_reply":"2024-07-23T12:19:32.136808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#helper function to set seed (from tutorials)\n\n# Call `set_seed` function in the exercises to ensure reproducibility.\nimport random\nimport torch\n\ndef set_seed(seed=None, seed_torch=True):\n  \"\"\"\n  Function that controls randomness. NumPy and random modules must be imported.\n\n  Args:\n    seed : Integer\n      A non-negative integer that defines the random state. Default is `None`.\n    seed_torch : Boolean\n      If `True` sets the random seed for pytorch tensors, so pytorch module\n      must be imported. Default is `True`.\n\n  Returns:\n    Nothing.\n  \"\"\"\n  if seed is None:\n    seed = np.random.choice(2 ** 32)\n  random.seed(seed)\n  np.random.seed(seed)\n  if seed_torch:\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.benchmark = False\n    torch.backends.cudnn.deterministic = True\n\n  print(f'Random seed {seed} has been set.')\n\n\n# In case that `DataLoader` is used\ndef seed_worker(worker_id):\n  \"\"\"\n  DataLoader will reseed workers following randomness in\n  multi-process data loading algorithm.\n\n  Args:\n    worker_id: integer\n      ID of subprocess to seed. 0 means that\n      the data will be loaded in the main process\n      Refer: https://pytorch.org/docs/stable/data.html#data-loading-randomness for more details\n\n  Returns:\n    Nothing\n  \"\"\"\n  worker_seed = torch.initial_seed() % 2**32\n  np.random.seed(worker_seed)\n  random.seed(worker_seed)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:19:35.448443Z","iopub.execute_input":"2024-07-23T12:19:35.448878Z","iopub.status.idle":"2024-07-23T12:19:35.459371Z","shell.execute_reply.started":"2024-07-23T12:19:35.448829Z","shell.execute_reply":"2024-07-23T12:19:35.458433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 2024\nset_seed(SEED)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:19:38.643216Z","iopub.execute_input":"2024-07-23T12:19:38.643574Z","iopub.status.idle":"2024-07-23T12:19:38.651937Z","shell.execute_reply.started":"2024-07-23T12:19:38.643542Z","shell.execute_reply":"2024-07-23T12:19:38.651051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!rm -rf /kaggle/working/*","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:19:40.658649Z","iopub.execute_input":"2024-07-23T12:19:40.659569Z","iopub.status.idle":"2024-07-23T12:19:40.663229Z","shell.execute_reply.started":"2024-07-23T12:19:40.659534Z","shell.execute_reply":"2024-07-23T12:19:40.662330Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## DATASET\n\nPrepares the dataset and saves it to kaggle disc - only need to run it once, if you have test_spec and train_spec files in '/kaggle/working' directory, just go to MODEL section","metadata":{}},{"cell_type":"code","source":"#load file with data description\nBASE_PATH = \"/kaggle/input/hms-harmful-brain-activity-classification\"\ndf = pd.read_csv(f'{BASE_PATH}/train.csv')\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:19:43.504487Z","iopub.execute_input":"2024-07-23T12:19:43.505315Z","iopub.status.idle":"2024-07-23T12:19:43.770499Z","shell.execute_reply.started":"2024-07-23T12:19:43.505283Z","shell.execute_reply":"2024-07-23T12:19:43.769559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SPEC_PATH = \"/kaggle/input/brain-eeg-spectrograms/EEG_Spectrograms\"\nos.chdir(SPEC_PATH)\nspec_list = os.listdir()\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"Spec_array = []\nlabels = []\npatient_ids = []\nfiles_to_add_later = {}\nlabs_to_add_later = {}\n\nfor spec_nme in spec_list:\n    spec = np.load(spec_nme)\n    \n    #find label\n    eegid = int(spec_nme.split('.')[0])\n    rec = df.loc[df['eeg_id'] == eegid]\n    lab = np.unique(rec['expert_consensus'].values)\n    #print(eegid, lab, len(lab))\n    \n    #patient id\n    pid = np.unique(rec['patient_id'].values)\n    \n    if len(lab) == 1 and lab != 'Other':\n        #only one file of each patient to train test split, then add remaining files by patient_id\n        if pid not in patient_ids:\n            Spec_array.append(spec)\n            labels.append(lab)\n            patient_ids.append(pid)\n        else:\n            if pid[0] not in list(files_to_add_later.keys()):\n                files_to_add_later[pid[0]] = spec\n                labs_to_add_later[pid[0]] = lab\n            #else:\n            #    files_to_add_later[pid[0]] = [files_to_add_later[pid[0]], spec]\n            #    labs_to_add_later[pid[0]] = [lab, lab]\n                                              #'lab': np.array([files_to_add_later[pid[0]]['lab'],lab])}\n\n\n        \n    \nnp.shape(Spec_array), np.shape(labels), len(files_to_add_later)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#classes - count occurences for each class\nnp.unique(labels, return_counts = True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#patients count\nids, num = np.unique(patient_ids, return_counts = True)\n#sum(num > 1)\nprint(len(ids))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#labels to int\nlabel_dict = {\n    'Seizure':1,\n    'LPD':2, #lateralized periodic discharges\n    'GPD':3,  #generalized periodic discharges, \n    'LRDA':4, #lateralized rhythmic delta activity\n    'GRDA':5 #, #generalized rhythmic delta activity\n    #'Other':0\n}\n\nlabels2 = [label_dict[l[0]] for l in labels]\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n_classes = len(np.unique(labels2))\nprint(n_classes)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#data loader\n\n#change array to tensor\nspec_data = torch.Tensor(np.array(Spec_array))\nlabels = torch.Tensor(np.array(labels2))\npatients = torch.Tensor(np.array(patient_ids))\n\n#set dimensions? (we have samples x freq x time x channels, if we need form samples x channels x time x freq)\nspec_data = spec_data.permute(0, 3, 1, 2)\n\n#train test split\ntrain_spec, test_spec, train_labels, test_labels, train_patients, test_patients = train_test_split(\n    spec_data, labels, patients, test_size=0.2, random_state=2024, shuffle=True, stratify=labels)\n\n#create dataset\n#train_spec_dataset = TensorDataset(train_spec, train_labels, train_patients)\n#my_dataloader = DataLoader(spec_dataset, batch_size=64, shuffle=True)\n\n#save dataset\n#torch.save(spec_dataset, '/kaggle/working/dataset_file') #- path??? and make it permanent between sessions?","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#add other files from train and test persons\ntrain_data_add = []\ntrain_labels_add = []\ntrain_pids_add = []\n\n#add files for training\nfor p in range(0, len(train_patients)):\n    pid = int(train_patients[p])\n    if pid in files_to_add_later.keys():\n        if len(files_to_add_later[pid]) == 128:\n            train_data_add.append(files_to_add_later[pid])\n            train_labels_add.append(labs_to_add_later[pid])\n            train_pids_add.append(pid)\n        #else:\n        #    for s in range (0, len(files_to_add_later[pid])):\n        #        train_data_add.append(files_to_add_later[pid][s])\n        #        train_labels_add.append(labs_to_add_later[pid][s])\n        #        train_pids_add.append(pid)\n            \ntrain_labels_add2 = [label_dict[l[0]] for l in train_labels_add]\n\nprint(np.shape(train_data_add))\n            \ntrain_spec_add = torch.Tensor(np.array(train_data_add))\ntrain_labs_add = torch.Tensor(np.array(train_labels_add2))\ntrain_patients_add = torch.Tensor(np.array(train_pids_add))\n\ntrain_spec_add = train_spec_add.permute(0, 3, 1, 2)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#the same for test data\ntest_data_add = []\ntest_labels_add = []\ntest_pids_add = []\n\n#add files for testing\nfor p in range(0, len(test_patients)):\n    pid = int(test_patients[p])\n    if pid in files_to_add_later.keys():\n        if len(files_to_add_later[pid]) == 128:\n            test_data_add.append(files_to_add_later[pid])\n            test_labels_add.append(labs_to_add_later[pid])\n            test_pids_add.append(pid)\n\n            \ntest_labels_add2 = [label_dict[l[0]] for l in test_labels_add]\n\nprint(np.shape(test_data_add))\n            \ntest_spec_add = torch.Tensor(np.array(test_data_add))\ntest_labs_add = torch.Tensor(np.array(test_labels_add2))\ntest_patients_add = torch.Tensor(np.array(test_pids_add))\n\ntest_spec_add = test_spec_add.permute(0, 3, 1, 2)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_patients.size()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_patients = torch.squeeze(train_patients)\ntest_patients = torch.squeeze(test_patients)\nprint(train_patients.size(), test_patients.size())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#join the tensors\ntest_spec = torch.cat((test_spec, test_spec_add), 0)\ntrain_spec = torch.cat((train_spec, train_spec_add), 0)\n\ntest_labels = torch.cat((test_labels, test_labs_add), 0)\ntrain_labels = torch.cat((train_labels, train_labs_add), 0)\n\ntest_patients = torch.cat((test_patients, test_patients_add), 0)\ntrain_patients = torch.cat((train_patients, train_patients_add), 0)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#calculate mean and std for normalization (only for train set)\nchann_mean = torch.mean(train_spec, axis=tuple([0, 2,3]))\nchann_sd = torch.std(train_spec, axis=tuple([0, 2,3]))\n\n#preprocessing transformations - normalization only?\npreproc_trans = v2.Compose([v2.Normalize(mean=chann_mean, std=chann_sd)]) #can add sth here","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#perform normalization\ntrain_spec_n = preproc_trans(train_spec)\ntest_spec_n = preproc_trans(test_spec)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"(train_spec_n.size(), train_labels.size())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#create and save dataset\ntrain_spec_dataset = TensorDataset(train_spec_n, train_labels) #, train_patients)\ntest_spec_dataset = TensorDataset(test_spec_n, test_labels) #, test_patients)\n\n#torch.save(train_spec_dataset, '/kaggle/working/train_spec.pt')\n#torch.save(test_spec_dataset, '/kaggle/working/test_spec.pt')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!find /kaggle/working -name \"*.pt\" -type f | zip kaggle_pth_files.zip -@","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndatasets_dict = {'train_data': train_spec_n, 'train_labels': train_labels, 'train_patients':train_patients,\n                'test_data': test_spec_n, 'test_labels': test_labels, 'test_patients':test_patients}\n#torch.save(datasets_dict, '/kaggle/working/data_specs_dict.pt')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#data augmentation / other transformations - think which would make sense\n#aug_trans = v2.Compose([...\n\n#all_trans = v2.Compose([preproc_trans, aug_trans])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## MODEL","metadata":{}},{"cell_type":"markdown","source":"Input size: 4, 128, 256, (channels x freq x time)","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\nimport torch.nn.functional as F","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:20:03.331513Z","iopub.execute_input":"2024-07-23T12:20:03.331906Z","iopub.status.idle":"2024-07-23T12:20:03.336911Z","shell.execute_reply.started":"2024-07-23T12:20:03.331851Z","shell.execute_reply":"2024-07-23T12:20:03.335795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_spec = torch.load('/kaggle/working/train_spec.pt')\ntest_spec = torch.load('/kaggle/working/test_spec.pt')","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:20:06.019680Z","iopub.execute_input":"2024-07-23T12:20:06.020077Z","iopub.status.idle":"2024-07-23T12:20:06.899917Z","shell.execute_reply.started":"2024-07-23T12:20:06.020044Z","shell.execute_reply":"2024-07-23T12:20:06.898908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEVICE = \"cuda\"\n#DEVICE = \"cpu\"","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:20:11.531300Z","iopub.execute_input":"2024-07-23T12:20:11.531650Z","iopub.status.idle":"2024-07-23T12:20:11.535691Z","shell.execute_reply.started":"2024-07-23T12:20:11.531621Z","shell.execute_reply":"2024-07-23T12:20:11.534731Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#train and test functions from the tutorials\ndef train(model, device, train_loader, epochs, lr2):\n  \"\"\"\n  Training function\n\n  Args:\n    model: nn.module\n      Neural network instance\n    device: string\n      GPU/CUDA if available, CPU otherwise\n    epochs: int\n      Number of epochs\n    train_loader: torch.loader\n      Training Set\n\n  Returns:\n    Nothing\n  \"\"\"\n  model.train()\n\n  loss_list = []\n  criterion = nn.CrossEntropyLoss()\n  optimizer = torch.optim.SGD(model.parameters(), lr=lr2)\n  for epoch in range(epochs):\n    loss_list_e = []\n    with tqdm(train_loader, unit='batch') as tepoch:\n      for data, target in tepoch:\n        data, target = data.to(device), target.to(device)\n        optimizer.zero_grad()\n        output = model(data)\n\n        loss = criterion(output, target.long())\n        loss.backward()\n        optimizer.step()\n        tepoch.set_postfix(loss=loss.item())\n        time.sleep(0.1)\n        loss_list_e.append(loss.item())\n    \n    loss_list.append(np.mean(loss_list_e))\n  return loss_list\n\ndef test(model, device, data_loader):\n  \"\"\"\n  Test function\n\n  Args:\n    model: nn.module\n      Neural network instance\n    device: string\n      GPU/CUDA if available, CPU otherwise\n    data_loader: torch.loader\n      Test Set\n\n  Returns:\n    acc: float\n      Test accuracy\n  \"\"\"\n  model.eval()\n  correct = 0\n  total = 0\n  for data in data_loader:\n    inputs, labels = data\n    inputs = inputs.to(device).float()\n    labels = labels.to(device).long()\n\n    outputs = model(inputs)\n    _, predicted = torch.max(outputs, 1)\n    total += labels.size(0)\n    correct += (predicted == labels).sum().item()\n\n  acc = 100 * correct / total\n  return acc","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:20:15.903484Z","iopub.execute_input":"2024-07-23T12:20:15.904337Z","iopub.status.idle":"2024-07-23T12:20:15.916322Z","shell.execute_reply.started":"2024-07-23T12:20:15.904302Z","shell.execute_reply":"2024-07-23T12:20:15.915266Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def calculate_output_shape(image_shape, kernel_shape):\n  \"\"\"\n  Helper function to calculate output shape\n\n  Args:\n    image_shape: tuple\n      Image shape\n    kernel_shape: tuple\n      Kernel shape\n\n  Returns:\n    output_height: int\n      Output Height\n    output_width: int\n      Output Width\n  \"\"\"\n  image_height, image_width = image_shape\n  kernel_height, kernel_width = kernel_shape\n  output_height = image_height - kernel_height + 1\n  output_width = image_width - kernel_width + 1\n  return output_height, output_width","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:20:21.147182Z","iopub.execute_input":"2024-07-23T12:20:21.147548Z","iopub.status.idle":"2024-07-23T12:20:21.153243Z","shell.execute_reply.started":"2024-07-23T12:20:21.147519Z","shell.execute_reply":"2024-07-23T12:20:21.152267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def trainN(model, device, train_loader, num_epochs, lr2, es):\n    # Track metrics\n    train_loss_list = []\n    val_loss_list = []\n    epochs = []\n    #test_loader = DataLoader(test_spec)\n\n    #num_epochs = 20  \n    early_stopping_patience = es  # Stop after 3 epochs with no improvement\n    best_val_loss = float('inf')\n    patience_counter = 0\n    \n    #optimizer, loss\n    #optimizer = torch.optim.SGD(model.parameters(), lr=lr2)\n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.Adam(model.parameters(), lr=lr2, weight_decay=1e-4)\n\n    for epoch in range(num_epochs):\n        model.train()\n        train_loss = 0.0\n\n        for inputs, labels in train_loader:\n            inputs, labels = inputs.to(device), labels.to(device)\n            optimizer.zero_grad()\n            outputs = model(inputs)\n            loss = criterion(outputs, labels.long())\n            loss.backward()\n            optimizer.step()\n            train_loss += loss.item() * inputs.size(0)\n\n        train_loss = train_loss / len(train_loader.dataset)\n        train_loss_list.append(train_loss)\n\n        model.eval()\n        val_loss = 0.0\n\n        with torch.no_grad():\n            for inputs, labels in test_loader:\n                inputs, labels = inputs.to(device), labels.to(device)\n                outputs = model(inputs)\n                loss = criterion(outputs, labels.long())\n                val_loss += loss.item() * inputs.size(0)\n\n        val_loss = val_loss / len(test_loader.dataset)\n        val_loss_list.append(val_loss)\n        epochs.append(epoch + 1)\n\n        print(f'Epoch [{epoch+1}/{num_epochs}], Train Loss: {train_loss:.4f}, Validation Loss: {val_loss:.4f}')\n\n        # Early stopping\n        #if val_loss < best_val_loss:\n        #    best_val_loss = val_loss\n        if train_loss < best_val_loss:\n            best_val_loss = train_loss\n            patience_counter = 0\n        else:\n            patience_counter += 1\n            if patience_counter >= early_stopping_patience:\n                print(\"Early stopping triggered\")\n                break\n\n    # Plot training and validation loss\n    plt.figure(figsize=(12, 6))\n    plt.subplot(1, 2, 1)\n    plt.plot(epochs, train_loss_list, 'r', label='Training loss')\n    plt.plot(epochs, val_loss_list, 'b', label='Validation loss')\n    plt.title('Training and Validation loss')\n    plt.xlabel('Epochs')\n    plt.ylabel('Loss')\n    plt.legend()\n    plt.tight_layout()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:20:28.544797Z","iopub.execute_input":"2024-07-23T12:20:28.545574Z","iopub.status.idle":"2024-07-23T12:20:28.559700Z","shell.execute_reply.started":"2024-07-23T12:20:28.545533Z","shell.execute_reply":"2024-07-23T12:20:28.558737Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"s1 = calculate_output_shape((256,128), (4,4))\ns2 = calculate_output_shape(s1, (2,2)) \ns3 = calculate_output_shape(s2, (2,2))\nprint(s1, s2, s3)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:20:33.972807Z","iopub.execute_input":"2024-07-23T12:20:33.973212Z","iopub.status.idle":"2024-07-23T12:20:33.979370Z","shell.execute_reply.started":"2024-07-23T12:20:33.973182Z","shell.execute_reply":"2024-07-23T12:20:33.978291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n_classes = 5","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:20:36.437078Z","iopub.execute_input":"2024-07-23T12:20:36.437461Z","iopub.status.idle":"2024-07-23T12:20:36.442221Z","shell.execute_reply.started":"2024-07-23T12:20:36.437429Z","shell.execute_reply":"2024-07-23T12:20:36.441233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SPEC_Net(nn.Module):\n\n\n  def __init__(self):\n\n    super(SPEC_Net, self).__init__()\n    self.conv1 = nn.Conv2d(in_channels=4, out_channels = 256, kernel_size=5, stride = 2)\n    self.conv2 = nn.Conv2d(256, 64, 4)\n    self.conv3 = nn.Conv2d(64, 32, 2)\n    self.conv4 = nn.Conv2d(32, 16, 2)\n    self.fc1 = nn.Linear(6032, 128)#(14880, 128)\n    self.fc2 = nn.Linear(128, n_classes)\n    self.pool = nn.MaxPool2d(2, 2)\n    self.dropout = nn.Dropout(0.2)\n    self.softmax = nn.Softmax(dim=1) #solution from net for training error?\n\n  def forward(self, x):\n    \"\"\"\n    Forward pass of SPECNet\n\n    Args:\n      x: torch.tensor\n        Input features\n\n    Returns:\n      x: torch.tensor\n        Output of final fully connected layer\n    \"\"\"\n    x = self.conv1(x)\n    x = F.relu(x)\n    x = self.conv2(x)\n    x = F.relu(x)\n    x = self.pool(x)\n    x = self.conv3(x)\n    x = F.relu(x)\n    x = self.pool(x)\n    x = self.conv4(x)\n    x = F.relu(x)\n    x = torch.flatten(x, 1)\n    #print(x[0].size())\n    x = self.fc1(x)\n    x = F.relu(x)\n    x = self.dropout(x)\n    x = self.fc2(x)\n    x = self.softmax(x)\n    #print(x)\n    return x\n\n\n\n\nspec_net = SPEC_Net().to(DEVICE)\nprint(\"Total Parameters in Network {:10d}\".format(sum(p.numel() for p in spec_net.parameters())))","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:24:13.751724Z","iopub.execute_input":"2024-07-23T12:24:13.752574Z","iopub.status.idle":"2024-07-23T12:24:13.781232Z","shell.execute_reply.started":"2024-07-23T12:24:13.752541Z","shell.execute_reply":"2024-07-23T12:24:13.780301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!find /kaggle/working -name \"*.pt\" -type f | zip kaggle_pth_files.zip -@","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#convert labels to 0, 1, 2, 3, 4\nnp.unique(np.array(train_spec[:][1]), return_counts = True)\n\ntrain_spec_l2 = torch.zeros_like(train_spec[:][1])\n\nfor l in range(0, len(train_spec[:][1])):\n    train_spec_l2[l] = train_spec[l][1] - torch.Tensor([1])\n    \ntrain_spec2 = TensorDataset(train_spec[:][0], train_spec_l2)\n\nnp.unique(np.array(train_spec2[:][1]), return_counts = True)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:24:19.831253Z","iopub.execute_input":"2024-07-23T12:24:19.832022Z","iopub.status.idle":"2024-07-23T12:24:19.914111Z","shell.execute_reply.started":"2024-07-23T12:24:19.831989Z","shell.execute_reply":"2024-07-23T12:24:19.913126Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_spec_l2 = torch.zeros_like(test_spec[:][1])\n\nfor l in range(0, len(test_spec[:][1])):\n    test_spec_l2[l] = test_spec[l][1] - torch.Tensor([1])\n    \ntest_spec2 = TensorDataset(test_spec[:][0], test_spec_l2)\n\nnp.unique(np.array(test_spec2[:][1]), return_counts = True)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:24:22.221004Z","iopub.execute_input":"2024-07-23T12:24:22.221647Z","iopub.status.idle":"2024-07-23T12:24:22.247789Z","shell.execute_reply.started":"2024-07-23T12:24:22.221615Z","shell.execute_reply":"2024-07-23T12:24:22.246798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#undersampling\n#pip install torchsampler\n!git clone https://github.com/ufoym/imbalanced-dataset-sampler.git\n!pip install torchsampler\n\nfrom torchsampler import ImbalancedDatasetSampler\n\n#train_loader = torch.utils.data.DataLoader(\n#    train_dataset,\n#    sampler=ImbalancedDatasetSampler(train_dataset),\n#    batch_size=args.batch_size,\n#    **kwargs\n#)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:24:25.734307Z","iopub.execute_input":"2024-07-23T12:24:25.735080Z","iopub.status.idle":"2024-07-23T12:24:41.437305Z","shell.execute_reply.started":"2024-07-23T12:24:25.735047Z","shell.execute_reply":"2024-07-23T12:24:41.435808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader = DataLoader(\n    train_spec2,\n    sampler=ImbalancedDatasetSampler(train_spec2),\n    batch_size=56)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:28:52.069824Z","iopub.execute_input":"2024-07-23T12:28:52.070290Z","iopub.status.idle":"2024-07-23T12:28:52.097461Z","shell.execute_reply.started":"2024-07-23T12:28:52.070243Z","shell.execute_reply":"2024-07-23T12:28:52.096309Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n_epochs = 50","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:28:54.515419Z","iopub.execute_input":"2024-07-23T12:28:54.515786Z","iopub.status.idle":"2024-07-23T12:28:54.520468Z","shell.execute_reply.started":"2024-07-23T12:28:54.515759Z","shell.execute_reply":"2024-07-23T12:28:54.519315Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#model, device, train_loader, num_epochs, lr2, es\nspec_net = SPEC_Net().to(DEVICE)\ntest_loader = DataLoader(test_spec2, sampler=ImbalancedDatasetSampler(test_spec2))\ntrainN(spec_net, DEVICE, train_loader, 50, 0.008, 15) #lr, early stop last 2 params","metadata":{"execution":{"iopub.status.busy":"2024-07-23T14:01:08.415687Z","iopub.execute_input":"2024-07-23T14:01:08.416734Z","iopub.status.idle":"2024-07-23T14:04:49.194015Z","shell.execute_reply.started":"2024-07-23T14:01:08.416689Z","shell.execute_reply":"2024-07-23T14:04:49.192938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"spec_net = SPEC_Net().to(DEVICE)\nloss_e = train(spec_net, DEVICE, train_loader, n_epochs, 0.008)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T09:26:34.151658Z","iopub.execute_input":"2024-07-23T09:26:34.152035Z","iopub.status.idle":"2024-07-23T09:34:02.437742Z","shell.execute_reply.started":"2024-07-23T09:26:34.152000Z","shell.execute_reply":"2024-07-23T09:34:02.436758Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#plot loss\nepochs = range(0, 50)\nplt.plot(epochs, loss_e)\n#batchsize 56, lr 0.008","metadata":{"execution":{"iopub.status.busy":"2024-07-23T09:35:21.896494Z","iopub.execute_input":"2024-07-23T09:35:21.896864Z","iopub.status.idle":"2024-07-23T09:35:22.176364Z","shell.execute_reply.started":"2024-07-23T09:35:21.896836Z","shell.execute_reply":"2024-07-23T09:35:22.175401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_loader = DataLoader(test_spec2, sampler=ImbalancedDatasetSampler(test_spec2))\ntest(spec_net, DEVICE, test_loader)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T09:35:38.953807Z","iopub.execute_input":"2024-07-23T09:35:38.954656Z","iopub.status.idle":"2024-07-23T09:35:41.182642Z","shell.execute_reply.started":"2024-07-23T09:35:38.954625Z","shell.execute_reply":"2024-07-23T09:35:41.181683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#for data in DataLoader(test_spec):\n#    inputs, labels = data\n#    inputs = inputs.to(DEVICE).float()\n#    labels = labels.to(DEVICE).long()#\n\n#    outputs = spec_net(inputs)\n#    _, predicted = torch.max(outputs, 1)\n#    #print(outputs)\n#    print(predicted, labels)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#torch.max(outputs, 1)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#TASK: IMPLEMENT PRE-TRAINED NETWORK","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Pretrained net\n\nSection in early progress...","metadata":{}},{"cell_type":"code","source":"","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:29:59.683627Z","iopub.execute_input":"2024-07-23T12:29:59.684037Z","iopub.status.idle":"2024-07-23T12:29:59.688651Z","shell.execute_reply.started":"2024-07-23T12:29:59.684003Z","shell.execute_reply":"2024-07-23T12:29:59.687538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"resnet = torchvision.models.resnet18(weights='ResNet18_Weights.DEFAULT')\nnum_ftrs = resnet.fc.in_features\n# Reset final fully connected layer, number of classes = types of Pokemon = 9\nresnet.fc = nn.Linear(num_ftrs, n_classes)\nresnet.to(DEVICE)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:30:05.318757Z","iopub.execute_input":"2024-07-23T12:30:05.319140Z","iopub.status.idle":"2024-07-23T12:30:06.099215Z","shell.execute_reply.started":"2024-07-23T12:30:05.319111Z","shell.execute_reply":"2024-07-23T12:30:06.098236Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"alexnet = torchvision.models.alexnet(weights='AlexNet_Weights.DEFAULT')\n#num_ftrs = alexnet.fc.in_features\n# Reset final fully connected layer, number of classes = types of Pokemon = 9\n#alexnet.fc = nn.Linear(num_ftrs, n_classes)\n#alexnet.to(DEVICE)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T13:59:46.492008Z","iopub.execute_input":"2024-07-23T13:59:46.492671Z","iopub.status.idle":"2024-07-23T13:59:47.384019Z","shell.execute_reply.started":"2024-07-23T13:59:46.492635Z","shell.execute_reply":"2024-07-23T13:59:47.382773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#alexnet.","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#transforms for alex net:\npreprocess_pretr = v2.Compose([\n                                 v2.Resize(256),\n                                 v2.CenterCrop(254),\n                                 #v2.ToTensor(),\n                                 #transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                                 #                     std=[0.229, 0.224, 0.225]),\n                                 ])","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:43:20.932318Z","iopub.execute_input":"2024-07-23T12:43:20.932791Z","iopub.status.idle":"2024-07-23T12:43:20.940837Z","shell.execute_reply.started":"2024-07-23T12:43:20.932752Z","shell.execute_reply":"2024-07-23T12:43:20.939982Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_spec_tr[:][1:4].size()","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:53:17.965368Z","iopub.execute_input":"2024-07-23T12:53:17.965787Z","iopub.status.idle":"2024-07-23T12:53:17.973251Z","shell.execute_reply.started":"2024-07-23T12:53:17.965753Z","shell.execute_reply":"2024-07-23T12:53:17.971960Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"execution":{"iopub.status.busy":"2024-07-23T13:00:21.277716Z","iopub.execute_input":"2024-07-23T13:00:21.278153Z","iopub.status.idle":"2024-07-23T13:00:21.284837Z","shell.execute_reply.started":"2024-07-23T13:00:21.278122Z","shell.execute_reply":"2024-07-23T13:00:21.283926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#dataset for resnet / alexnet:\ntrain_spec_tr = preprocess_pretr(train_spec2[:][0])\ndataset_pretr = TensorDataset(train_spec_tr[:, 1:4], train_spec2[:][1]) #only 3 fit channels\ndataset_pretr[:][0].size()","metadata":{"execution":{"iopub.status.busy":"2024-07-23T13:00:01.886226Z","iopub.execute_input":"2024-07-23T13:00:01.887242Z","iopub.status.idle":"2024-07-23T13:00:07.615185Z","shell.execute_reply.started":"2024-07-23T13:00:01.887203Z","shell.execute_reply":"2024-07-23T13:00:07.614074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#the same for test set\ntest_spec_tr = preprocess_pretr(test_spec2[:][0])\ndataset_pretr_t = TensorDataset(test_spec_tr[:, 1:4], test_spec2[:][1]) #only 3 first channels\ndataset_pretr_t[:][0].size()","metadata":{"execution":{"iopub.status.busy":"2024-07-23T13:06:15.563248Z","iopub.execute_input":"2024-07-23T13:06:15.564310Z","iopub.status.idle":"2024-07-23T13:06:17.039274Z","shell.execute_reply.started":"2024-07-23T13:06:15.564270Z","shell.execute_reply":"2024-07-23T13:06:17.038094Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test = preprocess_pretr(train_spec[9][0])","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:43:53.709677Z","iopub.execute_input":"2024-07-23T12:43:53.710751Z","iopub.status.idle":"2024-07-23T12:43:53.719665Z","shell.execute_reply.started":"2024-07-23T12:43:53.710714Z","shell.execute_reply":"2024-07-23T12:43:53.718770Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#do some function to lot it one one graph\n#fig, ax = plt.subplot etc\nsns.heatmap(train_spec[9][0][0])\nsns.heatmap(test[0])","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:45:05.713802Z","iopub.execute_input":"2024-07-23T12:45:05.714489Z","iopub.status.idle":"2024-07-23T12:45:06.521193Z","shell.execute_reply.started":"2024-07-23T12:45:05.714456Z","shell.execute_reply":"2024-07-23T12:45:06.520060Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sns.heatmap(test[0])","metadata":{"execution":{"iopub.status.busy":"2024-07-23T12:45:23.305592Z","iopub.execute_input":"2024-07-23T12:45:23.306354Z","iopub.status.idle":"2024-07-23T12:45:24.157542Z","shell.execute_reply.started":"2024-07-23T12:45:23.306320Z","shell.execute_reply":"2024-07-23T12:45:24.156504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pretrain_loader =  DataLoader(\n    dataset_pretr,\n    sampler=ImbalancedDatasetSampler(dataset_pretr),\n    batch_size=56)\n\npretrain_loader_test =  DataLoader(\n    dataset_pretr_t,\n    sampler=ImbalancedDatasetSampler(dataset_pretr_t))","metadata":{"execution":{"iopub.status.busy":"2024-07-23T13:06:54.118504Z","iopub.execute_input":"2024-07-23T13:06:54.118946Z","iopub.status.idle":"2024-07-23T13:06:54.135244Z","shell.execute_reply.started":"2024-07-23T13:06:54.118910Z","shell.execute_reply":"2024-07-23T13:06:54.134158Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(dataset_pretr)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T13:04:34.722833Z","iopub.execute_input":"2024-07-23T13:04:34.723724Z","iopub.status.idle":"2024-07-23T13:04:34.730141Z","shell.execute_reply.started":"2024-07-23T13:04:34.723690Z","shell.execute_reply":"2024-07-23T13:04:34.729128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_losses_pretr = []\ntrain_losses_pretr = []\n\n#optimizer = torch.optim.Adam(resnet.parameters(), lr=1e-4)\noptimizer = torch.optim.Adam(alexnet.parameters(), lr=1e-4)\nloss_fn = nn.CrossEntropyLoss()\nn_epochs = 50\n\n# @title Finetune ResNet\npretrained_accs = []\nfor epoch in tqdm(range(n_epochs)):\n  train_losses_pretr_b = []\n  # Train loop\n  for batch in pretrain_loader:\n    images, labels = batch\n    images = images.to(DEVICE)\n    labels = labels.to(DEVICE)\n\n\n    optimizer.zero_grad()\n    #output = resnet(images)\n    output = alexnet(images)\n    loss = loss_fn(output, labels.long())\n    train_losses_pretr_b.append(loss.item())\n    loss.backward()\n    optimizer.step()\n    #print(output)\n  train_losses_pretr.append(np.mean(train_losses_pretr_b))  \n    # Eval loop\n  with torch.no_grad():\n    loss_sum = 0\n    total_correct = 0\n    total = len(dataset_pretr_t)\n    for batch in pretrain_loader_test:\n      images, labels = batch\n      images = images.to(DEVICE)\n      labels = labels.to(DEVICE)\n      #output = resnet(images)\n      output = alexnet(images)\n      loss = loss_fn(output, labels.long())\n      loss_sum += loss.item()\n      test_losses_pretr.append(loss.item)\n\n      predictions = torch.argmax(output, dim=1)\n\n      num_correct = torch.sum(predictions == labels)\n      total_correct += num_correct\n\n# Plot accuracy\n    pretrained_accs.append(total_correct.cpu() / total)\n    \nplt.plot(pretrained_accs)\nplt.xlabel('epoch')\nplt.ylabel('accuracy')\n#plt.title('Resnet prediction accuracy')\nplt.title('Alexnet prediction accuracy')\nplt.show()\nplt.close()\n\n ","metadata":{"execution":{"iopub.status.busy":"2024-07-23T14:04:49.515067Z","iopub.status.idle":"2024-07-23T14:04:49.515659Z","shell.execute_reply.started":"2024-07-23T14:04:49.515355Z","shell.execute_reply":"2024-07-23T14:04:49.515379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pretrained_accs[-1]","metadata":{"execution":{"iopub.status.busy":"2024-07-23T13:36:49.120998Z","iopub.execute_input":"2024-07-23T13:36:49.121374Z","iopub.status.idle":"2024-07-23T13:36:49.129400Z","shell.execute_reply.started":"2024-07-23T13:36:49.121347Z","shell.execute_reply":"2024-07-23T13:36:49.128362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"math.ceil(len(train_losses_pretr) / 56)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T13:39:33.651365Z","iopub.execute_input":"2024-07-23T13:39:33.651898Z","iopub.status.idle":"2024-07-23T13:39:33.658840Z","shell.execute_reply.started":"2024-07-23T13:39:33.651827Z","shell.execute_reply":"2024-07-23T13:39:33.657748Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"    epochs = range(0, n_epochs)\n    plt.figure(figsize=(12, 6))\n    plt.subplot(1, 2, 1)\n    plt.plot(range(0, (n_epochs+1)*math.ceil(len(train_losses_pretr) / 56)), train_losses_pretr, 'r', label='Resnet Training loss')\n    plt.plot(range(0, n_epochs), test_losses_pretr, 'b', label='Resnet Test loss')\n    plt.title('Training and Validation loss')\n    plt.xlabel('Epochs')\n    plt.ylabel('Loss')\n    plt.legend()\n    plt.tight_layout()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-07-23T13:40:29.518619Z","iopub.execute_input":"2024-07-23T13:40:29.519384Z","iopub.status.idle":"2024-07-23T13:40:29.834010Z","shell.execute_reply.started":"2024-07-23T13:40:29.519348Z","shell.execute_reply":"2024-07-23T13:40:29.832240Z"},"trusted":true},"execution_count":null,"outputs":[]}]}