{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport numpy as np\nimport glob\nimport PIL.Image as Image\nimport torch.utils.data as data\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\nfrom tqdm import tqdm\nfrom ipywidgets import interact, fixed\n\nPREFIX = '/kaggle/input/vesuvius-challenge/train/1/'\n\nplt.imshow(Image.open(PREFIX+\"ir.png\"), cmap=\"gray\")","metadata":{"_uuid":"91db1348-a896-4607-8686-f6c6df6419ed","_cell_guid":"ff3c9fb9-0c86-4acf-9162-c741c46e53a4","collapsed":false,"jupyter":{"outputs_hidden":false},"_kg_hide-output":false,"execution":{"iopub.status.busy":"2023-04-09T19:49:42.577327Z","iopub.execute_input":"2023-04-09T19:49:42.577653Z","iopub.status.idle":"2023-04-09T19:49:49.771097Z","shell.execute_reply.started":"2023-04-09T19:49:42.577623Z","shell.execute_reply":"2023-04-09T19:49:49.769893Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\nmask = np.array(Image.open(PREFIX+\"mask.png\").convert('1'))\nlabel = torch.from_numpy(np.array(Image.open(PREFIX+\"inklabels.png\"))).gt(0).float().to(device)\nfig, (ax1, ax2) = plt.subplots(1, 2)\nax1.set_title(\"mask.png\")\nax1.imshow(mask, cmap='gray')\nax2.set_title(\"inklabels.png\")\nax2.imshow(label.cpu(), cmap='gray')\nplt.show()","metadata":{"_uuid":"351fefa9-30dd-4e8d-bf3b-1aaa0fb33905","_cell_guid":"43a9acbe-f2c5-4976-b9ee-027e62c27a83","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-04-09T19:49:49.772855Z","iopub.execute_input":"2023-04-09T19:49:49.773494Z","iopub.status.idle":"2023-04-09T19:49:58.107299Z","shell.execute_reply.started":"2023-04-09T19:49:49.773456Z","shell.execute_reply":"2023-04-09T19:49:58.106039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"---","metadata":{}},{"cell_type":"code","source":"import matplotlib.image as mpimg\nimport matplotlib.pyplot as plt\n\nimg=mpimg.imread('/kaggle/input/vesuvius-challenge/train/1/surface_volume/00.tif')\nimgplot = plt.imshow(img)","metadata":{"execution":{"iopub.status.busy":"2023-04-09T14:38:11.284520Z","iopub.execute_input":"2023-04-09T14:38:11.285298Z","iopub.status.idle":"2023-04-09T14:38:15.031632Z","shell.execute_reply.started":"2023-04-09T14:38:11.285258Z","shell.execute_reply":"2023-04-09T14:38:15.030522Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2\nfrom tqdm import tqdm\n\nres = []\nfor i in tqdm(range(65)):\n    if i < 10: i = f\"0{i}\"\n    img = cv2.imread(f'/kaggle/input/vesuvius-challenge/train/1/surface_volume/{i}.tif',\n                     cv2.IMREAD_UNCHANGED\n                    )\n    img = img / np.max(img) * 255.0\n    img = img.astype(\"uint8\")\n    res.append(img)\nres = np.array(res)\nnp.save(\"/kaggle/working/img1-uint8.npy\", res)","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:49:58.113123Z","iopub.execute_input":"2023-04-09T19:49:58.113786Z","iopub.status.idle":"2023-04-09T19:51:56.632524Z","shell.execute_reply.started":"2023-04-09T19:49:58.113743Z","shell.execute_reply":"2023-04-09T19:51:56.631265Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"res.shape","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:52:18.185970Z","iopub.execute_input":"2023-04-09T19:52:18.186440Z","iopub.status.idle":"2023-04-09T19:52:18.195311Z","shell.execute_reply.started":"2023-04-09T19:52:18.186387Z","shell.execute_reply":"2023-04-09T19:52:18.194011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"res = res[:2,:7500,:5500]","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:52:16.124168Z","iopub.execute_input":"2023-04-09T19:52:16.124606Z","iopub.status.idle":"2023-04-09T19:52:16.130286Z","shell.execute_reply.started":"2023-04-09T19:52:16.124569Z","shell.execute_reply":"2023-04-09T19:52:16.128901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\n\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:52:21.575803Z","iopub.execute_input":"2023-04-09T19:52:21.577001Z","iopub.status.idle":"2023-04-09T19:52:21.791340Z","shell.execute_reply.started":"2023-04-09T19:52:21.576953Z","shell.execute_reply":"2023-04-09T19:52:21.789594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(img[30,:,:], cmap='magma');","metadata":{"execution":{"iopub.status.busy":"2023-04-09T13:08:14.344235Z","iopub.execute_input":"2023-04-09T13:08:14.345061Z","iopub.status.idle":"2023-04-09T13:08:16.776283Z","shell.execute_reply.started":"2023-04-09T13:08:14.345020Z","shell.execute_reply":"2023-04-09T13:08:16.775286Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_np = label.cpu().detach().numpy()\nlabel = None","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:52:27.934172Z","iopub.execute_input":"2023-04-09T19:52:27.934597Z","iopub.status.idle":"2023-04-09T19:52:28.123198Z","shell.execute_reply.started":"2023-04-09T19:52:27.934560Z","shell.execute_reply":"2023-04-09T19:52:28.122121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"combined_mask = np.zeros_like(mask)\n\nmask[mask==1] = 1\nlabel_np[label_np==1] = 2\n\n# create target mask\ncombined_mask = mask + label_np\nplt.imshow(combined_mask);","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:52:29.712720Z","iopub.execute_input":"2023-04-09T19:52:29.713106Z","iopub.status.idle":"2023-04-09T19:52:33.114959Z","shell.execute_reply.started":"2023-04-09T19:52:29.713073Z","shell.execute_reply":"2023-04-09T19:52:33.113930Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask, label_np = None, None","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:52:33.116850Z","iopub.execute_input":"2023-04-09T19:52:33.118137Z","iopub.status.idle":"2023-04-09T19:52:33.138500Z","shell.execute_reply.started":"2023-04-09T19:52:33.118096Z","shell.execute_reply":"2023-04-09T19:52:33.137623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:52:33.849828Z","iopub.execute_input":"2023-04-09T19:52:33.851136Z","iopub.status.idle":"2023-04-09T19:52:34.029280Z","shell.execute_reply.started":"2023-04-09T19:52:33.851095Z","shell.execute_reply":"2023-04-09T19:52:34.028121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport torch\nfrom operator import truediv\n\ndef evaluate_accuracy(data_iter, net, model_vit, loss, device):\n    acc_sum, n = 0.0, 0\n    with torch.no_grad():\n        for X, y in data_iter:\n            test_l_sum, test_num = 0, 0\n            #X = X.permute(0, 3, 1, 2)\n            X = X.to(device)\n            y = y.to(device)\n            net.eval()\n            model_vit.eval()\n            y_hat = net(X)\n            y_hat = model_vit(y_hat)\n            l = loss(y_hat, y.long())\n            acc_sum += (y_hat.argmax(dim=1) == y.to(device)).float().sum().cpu().item()\n            test_l_sum += l\n            test_num += 1\n            net.train()\n            model_vit.train()\n            n += y.shape[0]\n    return [acc_sum / n, test_l_sum] # / test_num]\n\n\ndef aa_and_each_accuracy(confusion_matrix):\n    list_diag = np.diag(confusion_matrix)\n    list_raw_sum = np.sum(confusion_matrix, axis=1)\n    each_acc = np.nan_to_num(truediv(list_diag, list_raw_sum))\n    average_acc = np.mean(each_acc)\n    return each_acc, average_acc\n\n\n\ndef record_output(oa_ae, aa_ae, kappa_ae, element_acc_ae, training_time_ae, testing_time_ae, confusion_matrix, path):\n    f = open(path, 'a')\n    sentence0 = 'OAs for each iteration are:' + str(oa_ae) + '\\n'\n    f.write(sentence0)\n    sentence1 = 'AAs for each iteration are:' + str(aa_ae) + '\\n'\n    f.write(sentence1)\n    sentence2 = 'KAPPAs for each iteration are:' + str(kappa_ae) + '\\n' + '\\n'\n    f.write(sentence2)\n    sentence3 = 'mean_OA ± std_OA is: ' + str(np.mean(oa_ae)) + ' ± ' + str(np.std(oa_ae)) + '\\n'\n    f.write(sentence3)\n    sentence4 = 'mean_AA ± std_AA is: ' + str(np.mean(aa_ae)) + ' ± ' + str(np.std(aa_ae)) + '\\n'\n    f.write(sentence4)\n    sentence5 = 'mean_KAPPA ± std_KAPPA is: ' + str(np.mean(kappa_ae)) + ' ± ' + str(np.std(kappa_ae)) + '\\n' + '\\n'\n    f.write(sentence5)\n    sentence6 = 'Total average Training time is: ' + str(np.sum(training_time_ae)) + '\\n'\n    f.write(sentence6)\n    sentence7 = 'Total average Testing time is: ' + str(np.sum(testing_time_ae)) + '\\n' + '\\n'\n    f.write(sentence7)\n    element_mean = np.mean(element_acc_ae, axis=0)\n    element_std = np.std(element_acc_ae, axis=0)\n    sentence8 = \"Mean of all elements in confusion matrix: \" + str(element_mean) + '\\n'\n    f.write(sentence8)\n    sentence9 = \"Standard deviation of all elements in confusion matrix: \" + str(element_std) + '\\n'\n    f.write(sentence9)\n    sentence10 = \"The diagonal Confusion matrix: \" + '\\n' + str(confusion_matrix) + '\\n'\n    f.write(sentence10)\n    f.close()\n","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:52:40.359503Z","iopub.execute_input":"2023-04-09T19:52:40.360064Z","iopub.status.idle":"2023-04-09T19:52:40.381545Z","shell.execute_reply.started":"2023-04-09T19:52:40.360010Z","shell.execute_reply":"2023-04-09T19:52:40.380099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport numpy as np\nimport torch.utils.data as Data\n\ndef index_assignment(index, row, col, pad_length):\n    new_assign = {}\n    for counter, value in enumerate(index):\n        assign_0 = value // col + pad_length\n        assign_1 = value % col + pad_length\n        new_assign[counter] = [assign_0, assign_1]\n        \n    gc.collect()\n    return new_assign\n\ndef select_patch(matrix, pos_row, pos_col, ex_len):\n    selected_rows = matrix[range(pos_row-ex_len, pos_row+ex_len+1)]\n    selected_patch = selected_rows[:, range(pos_col-ex_len, pos_col+ex_len+1)]\n    \n    gc.collect()\n    return selected_patch\n\n\ndef select_small_cubic(data_size, data_indices, whole_data, patch_length, padded_data, dimension):\n    small_cubic_data = np.zeros_like((data_size, 2 * patch_length + 1, 2 * patch_length + 1, dimension))\n    data_assign = index_assignment(data_indices, whole_data.shape[0], whole_data.shape[1], patch_length)\n    for i in range(len(data_assign)):\n        small_cubic_data[i] = select_patch(padded_data, data_assign[i][0], data_assign[i][1], patch_length)\n        \n    gc.collect()\n    return small_cubic_data\n\n\ndef generate_iter(TRAIN_SIZE, train_indices, TEST_SIZE, test_indices, TOTAL_SIZE, total_indices, TOTAL_SIZEBG, total_indicesbg, VAL_SIZE,\n                  whole_data, PATCH_LENGTH, padded_data, INPUT_DIMENSION, batch_size, gt):\n    gt_all_bg = gt[total_indicesbg]\n    gt_all = gt[total_indices] - 1\n    y_train = gt[train_indices] - 1\n    y_test = gt[test_indices] - 1\n\n    all_data =  select_small_cubic(TOTAL_SIZE, total_indices, whole_data,\n                                                      PATCH_LENGTH, padded_data, INPUT_DIMENSION)\n\n    all_data_bg =  select_small_cubic(TOTAL_SIZEBG, total_indicesbg, whole_data,\n                                                      PATCH_LENGTH, padded_data, INPUT_DIMENSION)\n\n    \n    train_data = select_small_cubic(TRAIN_SIZE, train_indices, whole_data,\n                                                        PATCH_LENGTH, padded_data, INPUT_DIMENSION)\n    test_data =  select_small_cubic(TEST_SIZE, test_indices, whole_data,\n                                                       PATCH_LENGTH, padded_data, INPUT_DIMENSION)\n    x_train = train_data.reshape(train_data.shape[0], train_data.shape[1], train_data.shape[2], INPUT_DIMENSION)\n    x_test_all = test_data.reshape(test_data.shape[0], test_data.shape[1], test_data.shape[2], INPUT_DIMENSION)\n\n    x_val = x_test_all[-VAL_SIZE:]\n    y_val = y_test[-VAL_SIZE:]\n\n    x_test = x_test_all[:-VAL_SIZE]\n    y_test = y_test[:-VAL_SIZE]\n    \n    gc.collect()\n    \n    x1_tensor_train = torch.from_numpy(x_train).type(torch.FloatTensor).unsqueeze(1)\n    y1_tensor_train = torch.from_numpy(y_train).type(torch.FloatTensor)\n    torch_dataset_train = Data.TensorDataset(x1_tensor_train, y1_tensor_train)\n\n    x1_tensor_valida = torch.from_numpy(x_val).type(torch.FloatTensor).unsqueeze(1)\n    y1_tensor_valida = torch.from_numpy(y_val).type(torch.FloatTensor)\n    torch_dataset_valida = Data.TensorDataset(x1_tensor_valida, y1_tensor_valida)\n\n    x1_tensor_test = torch.from_numpy(x_test).type(torch.FloatTensor).unsqueeze(1)\n    y1_tensor_test = torch.from_numpy(y_test).type(torch.FloatTensor)\n    torch_dataset_test = Data.TensorDataset(x1_tensor_test,y1_tensor_test)\n\n    all_data.reshape(all_data.shape[0], all_data.shape[1], all_data.shape[2], INPUT_DIMENSION)\n    all_tensor_data = torch.from_numpy(all_data).type(torch.FloatTensor).unsqueeze(1)\n    all_tensor_data_label = torch.from_numpy(gt_all).type(torch.FloatTensor)\n    torch_dataset_all = Data.TensorDataset(all_tensor_data, all_tensor_data_label)\n\n    all_data_bg.reshape(all_data_bg.shape[0], all_data_bg.shape[1], all_data_bg.shape[2], INPUT_DIMENSION)\n    all_tensor_data_bg = torch.from_numpy(all_data_bg).type(torch.FloatTensor).unsqueeze(1)\n    all_tensor_data_label_bg = torch.from_numpy(gt_all_bg).type(torch.FloatTensor)\n    torch_dataset_all_bg = Data.TensorDataset(all_tensor_data_bg, all_tensor_data_label_bg)\n    \n    gc.collect()\n\n    train_iter = Data.DataLoader(\n        dataset=torch_dataset_train,  # torch TensorDataset format\n        batch_size=batch_size,  # mini batch size\n        shuffle=True,  \n        num_workers=0, \n    )\n    valiada_iter = Data.DataLoader(\n        dataset=torch_dataset_valida,  # torch TensorDataset format\n        batch_size=batch_size,  # mini batch size\n        shuffle=True,  \n        num_workers=0, \n    )\n    test_iter = Data.DataLoader(\n        dataset=torch_dataset_test,  # torch TensorDataset format\n        batch_size=batch_size,  # mini batch size\n        shuffle=False, \n        num_workers=0, \n    )\n    all_iter = Data.DataLoader(\n        dataset=torch_dataset_all,  # torch TensorDataset format\n        batch_size=batch_size,  # mini batch size\n        shuffle=False, \n        num_workers=0, \n    )\n    all_iter_bg = Data.DataLoader(\n        dataset=torch_dataset_all_bg,  # torch TensorDataset format\n        batch_size=batch_size,  # mini batch size\n        shuffle=False, \n        num_workers=0, \n    )\n    \n    gc.collect()\n    \n    return train_iter, valiada_iter, test_iter, all_iter, all_iter_bg #, y_test","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:52:42.121340Z","iopub.execute_input":"2023-04-09T19:52:42.121863Z","iopub.status.idle":"2023-04-09T19:52:42.145113Z","shell.execute_reply.started":"2023-04-09T19:52:42.121828Z","shell.execute_reply":"2023-04-09T19:52:42.143651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nfrom sklearn import metrics, preprocessing\nfrom sklearn.preprocessing import MinMaxScaler\nfrom sklearn.decomposition import PCA\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import confusion_matrix, accuracy_score, classification_report, cohen_kappa_score\nfrom operator import truediv\nfrom plotly.offline import init_notebook_mode\nimport matplotlib.pyplot as plt\nimport scipy.io as sio\nimport os\nimport spectral\nimport torch\nimport cv2\nfrom operator import truediv\nfrom IPython import display\n\n\ndef sampling(proportion, ground_truth, bg=False):\n    train = {}\n    test = {}\n    labels_loc = {}\n    m = max(ground_truth)+1\n    a = 0 if bg == True else 1\n    for i in range(m):\n        indexes = [j for j, x in enumerate(ground_truth.ravel().tolist()) if x == i + a]\n        np.random.shuffle(indexes)\n        labels_loc[i] = indexes\n        if proportion != 1:\n            nb_val = max(int((1 - proportion) * len(indexes)), 3)\n        else:\n            nb_val = 0\n        train[i] = indexes[:nb_val]\n        test[i] = indexes[nb_val:]\n        \n        gc.collect()\n        \n    train_indexes = []\n    test_indexes = []\n    for i in range(m):\n        train_indexes += train[i]\n        test_indexes += test[i]\n    np.random.shuffle(train_indexes)\n    np.random.shuffle(test_indexes)\n    \n    gc.collect()\n    \n    return train_indexes, test_indexes\n\n\n\ndef set_figsize(figsize=(3.5, 2.5)):\n    display.set_matplotlib_formats('svg')\n    plt.rcParams['figure.figsize'] = figsize\n\n\ndef classification_map(map, ground_truth, dpi, save_path):\n    fig = plt.figure(frameon=False)\n    fig.set_size_inches(ground_truth.shape[1] * 2.0 / dpi, ground_truth.shape[0] * 2.0 / dpi)\n    ax = plt.Axes(fig, [0., 0., 1., 1.])\n    ax.set_axis_off()\n    ax.xaxis.set_visible(False)\n    ax.yaxis.set_visible(False)\n    fig.add_axes(ax)\n    ax.imshow(map)\n    fig.savefig(save_path, dpi=dpi)\n    return 0\n\n\ndef list_to_colormap(x_list):\n    y = np.zeros((x_list.shape[0], 3))\n    for index, item in enumerate(x_list):\n        if item == 0:\n            y[index] = np.array([255, 0, 0]) / 255.\n        if item == 1:\n            y[index] = np.array([0, 255, 0]) / 255.\n        if item == 2:\n            y[index] = np.array([0, 0, 255]) / 255.\n        if item == 3:\n            y[index] = np.array([255, 255, 0]) / 255.\n        if item == 4:\n            y[index] = np.array([0, 255, 255]) / 255.\n        if item == 5:\n            y[index] = np.array([255, 0, 255]) / 255.\n        if item == 6:\n            y[index] = np.array([192, 192, 192]) / 255.\n        if item == 7:\n            y[index] = np.array([128, 128, 128]) / 255.\n        if item == 8:\n            y[index] = np.array([128, 0, 0]) / 255.\n        if item == 9:\n            y[index] = np.array([128, 128, 0]) / 255.\n        if item == 10:\n            y[index] = np.array([0, 128, 0]) / 255.\n        if item == 11:\n            y[index] = np.array([128, 0, 128]) / 255.\n        if item == 12:\n            y[index] = np.array([0, 128, 128]) / 255.\n        if item == 13:\n            y[index] = np.array([0, 0, 128]) / 255.\n        if item == 14:\n            y[index] = np.array([255, 165, 0]) / 255.\n        if item == 15:\n            y[index] = np.array([255, 215, 0]) / 255.\n        if item == 16:\n            y[index] = np.array([0, 0, 0]) / 255.\n        if item == 17:\n            y[index] = np.array([215, 255, 0]) / 255.\n        if item == 18:\n            y[index] = np.array([0, 255, 215]) / 255.\n        if item == -1:\n            y[index] = np.array([0, 0, 0]) / 255.\n    return y\n\n\n\ndef generate_png(all_iter, net, model_vit, gt_hsi, Dataset, device, total_indices):\n    pred_test = []\n    for X, y in all_iter:\n        #X = X.permute(0, 3, 1, 2)\n        X = X.to(device)\n        net.eval()\n        model_vit.eval()\n        y_hat = net(X)\n        y_hat = model_vit(y_hat)\n        pred_test.extend(y_hat.cpu().argmax(axis=1).detach().numpy())\n    gt = gt_hsi.flatten()\n    \n    x_label = np.zeros(gt.shape)\n\n    gt = gt[:] - 1\n\n    x_label[total_indices] = pred_test\n    x = np.ravel(x_label)\n    y_list = list_to_colormap(x)\n\n    y_list2=[]\n    for xx,yy in zip(x,gt):\n        if yy == 255: y_list2.append(16)\n        else: y_list2.append(xx)\n    y_list2 = np.array(y_list2)\n\n    y_gt = list_to_colormap(gt)\n    y_re = np.reshape(y_list, (gt_hsi.shape[0], gt_hsi.shape[1], 3))\n    y_re2 = np.reshape(list_to_colormap(y_list2), (gt_hsi.shape[0], gt_hsi.shape[1], 3))\n    gt_re = np.reshape(y_gt, (gt_hsi.shape[0], gt_hsi.shape[1], 3))\n    path = './content/' \n    classification_map(y_re, gt_hsi, 300,\n                       path + '/classification_maps/' + Dataset + '_BG'  +  '.png')\n    classification_map(y_re2, gt_hsi, 300,\n                       path + '/classification_maps/' + Dataset + '_'  +  '.png')\n    classification_map(gt_re, gt_hsi, 300,\n                       path + '/classification_maps/' + Dataset + '_gt.png')\n    print('------Get classification maps successful-------')","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:52:47.063358Z","iopub.execute_input":"2023-04-09T19:52:47.064173Z","iopub.status.idle":"2023-04-09T19:52:47.930756Z","shell.execute_reply.started":"2023-04-09T19:52:47.064133Z","shell.execute_reply":"2023-04-09T19:52:47.929534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn import metrics, preprocessing\nfrom sklearn.preprocessing import MinMaxScaler\nfrom sklearn.decomposition import PCA\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import confusion_matrix, accuracy_score, classification_report, cohen_kappa_score\n\nfrom operator import truediv\n\nfrom plotly.offline import init_notebook_mode\n\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport scipy.io as sio\nimport os\nimport spectral\nimport cv2\nfrom operator import truediv\n\ninit_notebook_mode(connected=True)\n%matplotlib inline","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:52:52.796632Z","iopub.execute_input":"2023-04-09T19:52:52.797396Z","iopub.status.idle":"2023-04-09T19:52:52.888538Z","shell.execute_reply.started":"2023-04-09T19:52:52.797353Z","shell.execute_reply":"2023-04-09T19:52:52.887189Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# for Monte Carlo runs\nseeds = [1331, 1332, 1333, 1334, 1335, 1336, 1337, 1338, 1339, 1340, 1341]\nensemble = 1","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:52:55.700694Z","iopub.execute_input":"2023-04-09T19:52:55.701565Z","iopub.status.idle":"2023-04-09T19:52:55.707768Z","shell.execute_reply.started":"2023-04-09T19:52:55.701524Z","shell.execute_reply":"2023-04-09T19:52:55.706380Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"global Dataset  # UP,IN,SV\ndataset = '1' #input('Please input the name of Dataset(IN, UP, SV):')\nDataset = dataset.upper()","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:52:57.705885Z","iopub.execute_input":"2023-04-09T19:52:57.706246Z","iopub.status.idle":"2023-04-09T19:52:57.712679Z","shell.execute_reply.started":"2023-04-09T19:52:57.706213Z","shell.execute_reply":"2023-04-09T19:52:57.710614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_dataset(Dataset):\n        \n    if Dataset == '1':\n        data_hsi = res\n        gt_hsi = combined_mask.astype('int8')\n        K = 15 # 3\n        TOTAL_SIZE = res.shape[0] * res.shape[1] * res.shape[2]\n        VALIDATION_SPLIT = 0.90\n        TRAIN_SIZE = math.ceil(TOTAL_SIZE * VALIDATION_SPLIT)\n\n\n    #shapeor=data_hsi.shape\n    #data_hsi = data_hsi.reshape(-1,data_hsi.shape[-1])\n    #data_hsi = PCA(n_components=K).fit_transform(data_hsi)\n    #shapeor = np.array(shapeor)\n    #shapeor[-1] = K\n    #data_hsi = data_hsi.reshape(shapeor)    \n\n\n    return data_hsi, gt_hsi, TOTAL_SIZE, TRAIN_SIZE, VALIDATION_SPLIT\n","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:53:20.515272Z","iopub.execute_input":"2023-04-09T19:53:20.515638Z","iopub.status.idle":"2023-04-09T19:53:20.526067Z","shell.execute_reply.started":"2023-04-09T19:53:20.515604Z","shell.execute_reply":"2023-04-09T19:53:20.524769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\ndata_hsi, gt_hsi, TOTAL_SIZE, TRAIN_SIZE,VALIDATION_SPLIT = load_dataset(Dataset)\nprint(data_hsi.shape)\nimage_x, image_y, BAND = data_hsi.shape\ndata = data_hsi.reshape(np.prod(data_hsi.shape[:2]), np.prod(data_hsi.shape[2:]))\ngt = gt_hsi.reshape(np.prod(gt_hsi.shape[:2]),)\nCLASSES_NUM = max(gt).astype(\"int\")\nprint('The class numbers of the HSI data is:', CLASSES_NUM)\n\nprint('-----Importing Setting Parameters-----')\nITER = 1\nPATCH_LENGTH = 4\nlr, num_epochs, batch_size = 0.001, 100, 32\nloss = torch.nn.CrossEntropyLoss()\n\nimg_rows = 2*PATCH_LENGTH+1\nimg_cols = 2*PATCH_LENGTH+1\nimg_channels = data_hsi.shape[2]\nINPUT_DIMENSION = data_hsi.shape[2]\nALL_SIZE = data_hsi.shape[0] * data_hsi.shape[1]\nVAL_SIZE = int(TRAIN_SIZE)\nTEST_SIZE = TOTAL_SIZE - TRAIN_SIZE\n\n\nKAPPA = []\nOA = []\nAA = []\nTRAINING_TIME = []\nTESTING_TIME = []\nELEMENT_ACC = np.zeros_like((ITER, CLASSES_NUM))\n\ndata = preprocessing.scale(data)\ndata_ = data.reshape(data_hsi.shape[0], data_hsi.shape[1], data_hsi.shape[2])\nwhole_data = data_\npadded_data = np.lib.pad(whole_data, ((PATCH_LENGTH, PATCH_LENGTH), (PATCH_LENGTH, PATCH_LENGTH), (0, 0)),\n                         'constant', constant_values=0)","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:53:30.473787Z","iopub.execute_input":"2023-04-09T19:53:30.474485Z","iopub.status.idle":"2023-04-09T19:53:39.672705Z","shell.execute_reply.started":"2023-04-09T19:53:30.474446Z","shell.execute_reply":"2023-04-09T19:53:39.671584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.nn import init","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:53:39.674618Z","iopub.execute_input":"2023-04-09T19:53:39.674977Z","iopub.status.idle":"2023-04-09T19:53:39.680080Z","shell.execute_reply.started":"2023-04-09T19:53:39.674941Z","shell.execute_reply":"2023-04-09T19:53:39.678936Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Residual(nn.Module):  # pytorch\n    def __init__(self, in_channels, out_channels, kernel_size, padding, use_1x1conv=False, stride=1):\n        super(Residual, self).__init__()\n        self.conv1 = nn.Sequential(\n            nn.Conv3d(in_channels, out_channels,\n                      kernel_size=kernel_size, padding=padding, stride=stride),\n            nn.ReLU()\n        )\n        self.conv2 = nn.Conv3d(out_channels, out_channels,\n                               kernel_size=kernel_size, padding=padding,stride=stride)\n        if use_1x1conv:\n            self.conv3 = nn.Conv3d(in_channels, out_channels, kernel_size=1, stride=stride)\n        else:\n            self.conv3 = None\n        self.bn1 = nn.BatchNorm3d(out_channels)\n        self.bn2 = nn.BatchNorm3d(out_channels)\n\n    def forward(self, X):\n        Y = F.relu(self.bn1(self.conv1(X)))\n        Y = self.bn2(self.conv2(Y))\n        if self.conv3:\n            X = self.conv3(X)\n        return F.relu(Y + X)\n\nclass SSRN_network(nn.Module):\n    def __init__(self, band, classes):\n        super(SSRN_network, self).__init__()\n        self.name = 'SSRN'\n        self.conv1 = nn.Conv3d(in_channels=1, out_channels=24,\n                                kernel_size=(1, 1, 7), stride=(1, 1, 2))\n        self.batch_norm1 = nn.Sequential(\n            nn.BatchNorm3d(24, eps=0.001, momentum=0.1, affine=True),  # 0.1\n            nn.ReLU(inplace=True)\n        )\n\n        self.res_net1 = Residual(24, 24, (1, 1, 7), (0, 0, 3))\n        self.res_net2 = Residual(24, 24, (1, 1, 7), (0, 0, 3))\n        self.res_net3 = Residual(24, 24, (3, 3, 1), (1, 1, 0))\n        self.res_net4 = Residual(24, 24, (3, 3, 1), (1, 1, 0))\n\n        kernel_3d = math.ceil((band - 6) / 2)\n\n        self.conv2 = nn.Conv3d(in_channels=24, out_channels=128, padding=(0, 0, 0),\n                               kernel_size=(1, 1, kernel_3d), stride=(1, 1, 1))\n        self.batch_norm2 = nn.Sequential(\n            nn.BatchNorm3d(128, eps=0.001, momentum=0.1, affine=True),  # 0.1\n            nn.ReLU(inplace=True)\n        )\n        self.conv3 = nn.Conv3d(in_channels=1, out_channels=24, padding=(0, 0, 0),\n                               kernel_size=(3, 3, 128), stride=(1, 1, 1))\n        self.batch_norm3 = nn.Sequential(\n            nn.BatchNorm3d(24, eps=0.001, momentum=0.1, affine=True),  # 0.1\n            nn.ReLU(inplace=True)\n        )\n\n        self.avg_pooling = nn.AvgPool3d(kernel_size=(5, 5, 1))\n        self.full_connection = nn.Sequential(\n            # nn.Dropout(p=0.5),\n            nn.Linear(24, classes)  # ,\n            # nn.Softmax()\n        )\n\n    def forward(self, X):\n        x1 = self.batch_norm1(self.conv1(X))\n        # print('x1', x1.shape)\n\n        x2 = self.res_net3(x1)\n        x2 = self.res_net4(x2)\n        x2 = self.batch_norm2(self.conv2(x2))\n        x2 = x2.permute(0, 4, 2, 3, 1)\n        x2 = self.batch_norm3(self.conv3(x2))\n\n        x2 = x2.view(x2.size()[0],x2.size()[1]*x2.size()[4],x2.size()[2],x2.size()[3])\n\n        return x2","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:53:40.086764Z","iopub.execute_input":"2023-04-09T19:53:40.087272Z","iopub.status.idle":"2023-04-09T19:53:40.117656Z","shell.execute_reply.started":"2023-04-09T19:53:40.087233Z","shell.execute_reply":"2023-04-09T19:53:40.115072Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = SSRN_network(200, 16).cuda()","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:53:42.374406Z","iopub.execute_input":"2023-04-09T19:53:42.374771Z","iopub.status.idle":"2023-04-09T19:53:42.418290Z","shell.execute_reply.started":"2023-04-09T19:53:42.374736Z","shell.execute_reply":"2023-04-09T19:53:42.417186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install torch-summary","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:45:45.206090Z","iopub.execute_input":"2023-04-09T19:45:45.207246Z","iopub.status.idle":"2023-04-09T19:45:57.183747Z","shell.execute_reply.started":"2023-04-09T19:45:45.207194Z","shell.execute_reply":"2023-04-09T19:45:57.182488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchsummary import summary\n\nsummary(model, (1,9,9,200))","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:45:57.186499Z","iopub.execute_input":"2023-04-09T19:45:57.187451Z","iopub.status.idle":"2023-04-09T19:45:59.859272Z","shell.execute_reply.started":"2023-04-09T19:45:57.187401Z","shell.execute_reply":"2023-04-09T19:45:59.857961Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:53:47.089237Z","iopub.execute_input":"2023-04-09T19:53:47.089601Z","iopub.status.idle":"2023-04-09T19:53:47.291738Z","shell.execute_reply.started":"2023-04-09T19:53:47.089567Z","shell.execute_reply":"2023-04-09T19:53:47.290316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import ctypes\nlibc = ctypes.CDLL(\"libc.so.6\") # clearing cache \nlibc.malloc_trim(0)","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:53:48.834470Z","iopub.execute_input":"2023-04-09T19:53:48.834836Z","iopub.status.idle":"2023-04-09T19:53:48.842793Z","shell.execute_reply.started":"2023-04-09T19:53:48.834802Z","shell.execute_reply":"2023-04-09T19:53:48.841389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install einops\nimport torch\nimport torch.nn.functional as F\nfrom einops import rearrange, repeat\nfrom torch import nn\n\nMIN_NUM_PATCHES = 16\n\nclass Residual_1(nn.Module):\n    def __init__(self, fn):\n        super().__init__()\n        self.fn = fn\n    def forward(self, x, **kwargs):\n        return self.fn(x, **kwargs) + x\n\nclass PreNorm(nn.Module):\n    def __init__(self, dim, fn):\n        super().__init__()\n        self.norm = nn.LayerNorm(dim)\n        self.fn = fn\n    def forward(self, x, **kwargs):\n        return self.fn(self.norm(x), **kwargs)\n\nclass FeedForward(nn.Module):\n    def __init__(self, dim, hidden_dim, dropout = 0.):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Linear(dim, hidden_dim),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(hidden_dim, dim),\n            nn.Dropout(dropout)\n        )\n    def forward(self, x):\n        return self.net(x)\n\nclass Attention(nn.Module):\n    def __init__(self, dim, heads = 8, dim_head = 64, dropout = 0.):\n        super().__init__()\n        inner_dim = dim_head *  heads\n        self.heads = heads\n        self.scale = dim ** -0.5\n\n        self.to_qkv = nn.Linear(dim, inner_dim * 3, bias = False)\n        self.to_out = nn.Sequential(\n            nn.Linear(inner_dim, dim),\n            nn.Dropout(dropout)\n        )\n\n    def forward(self, x, mask = None):\n        b, n, _, h = *x.shape, self.heads\n        qkv = self.to_qkv(x).chunk(3, dim = -1)\n        q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> b h n d', h = h), qkv)\n\n        dots = torch.einsum('bhid,bhjd->bhij', q, k) * self.scale\n        mask_value = -torch.finfo(dots.dtype).max\n\n        if mask is not None:\n            mask = F.pad(mask.flatten(1), (1, 0), value = True)\n            assert mask.shape[-1] == dots.shape[-1], 'mask has incorrect dimensions'\n            mask = mask[:, None, :] * mask[:, :, None]\n            dots.masked_fill_(~mask, mask_value)\n            del mask\n\n        attn = dots.softmax(dim=-1)\n\n        out = torch.einsum('bhij,bhjd->bhid', attn, v)\n        out = rearrange(out, 'b h n d -> b n (h d)')\n        out =  self.to_out(out)\n        return out\n\nclass Transformer(nn.Module):\n    def __init__(self, dim, depth, heads, dim_head, mlp_dim, dropout):\n        super().__init__()\n        self.layers = nn.ModuleList([])\n        for _ in range(depth):\n            self.layers.append(nn.ModuleList([\n                Residual_1(PreNorm(dim, Attention(dim, heads = heads, dim_head = dim_head, dropout = dropout))),\n                Residual_1(PreNorm(dim, FeedForward(dim, mlp_dim, dropout = dropout)))\n            ]))\n    def forward(self, x, mask = None):\n        for attn, ff in self.layers:\n            x = attn(x, mask = mask)\n            x = ff(x)\n        return x\n\nclass ViT(nn.Module):\n    def __init__(self, *, image_size, patch_size, num_classes, dim, depth, heads, mlp_dim, channels = 3, dim_head = 64, dropout = 0., emb_dropout = 0.):\n        super().__init__()\n        assert image_size % patch_size == 0, 'Image dimensions must be divisible by the patch size.'\n        num_patches = (image_size // patch_size) ** 2\n        patch_dim = channels * patch_size ** 2\n        assert num_patches > MIN_NUM_PATCHES, f'your number of patches ({num_patches}) is way too small for attention to be effective (at least 16). Try decreasing your patch size'\n\n        self.patch_size = patch_size\n\n        self.pos_embedding = nn.Parameter(torch.randn(1, num_patches + 1, dim))\n        self.patch_to_embedding = nn.Linear(patch_dim, dim)\n        self.cls_token = nn.Parameter(torch.randn(1, 1, dim))\n        self.dropout = nn.Dropout(emb_dropout)\n\n        self.transformer = Transformer(dim, depth, heads, dim_head, mlp_dim, dropout)\n\n        self.to_cls_token = nn.Identity()\n\n        self.mlp_head = nn.Sequential(\n            nn.LayerNorm(dim),\n            nn.Linear(dim, num_classes)\n        )\n\n    def forward(self, img, mask = None):\n        p = self.patch_size\n\n        x = rearrange(img, 'b c (h p1) (w p2) -> b (h w) (p1 p2 c)', p1 = p, p2 = p)\n        x = self.patch_to_embedding(x)\n        b, n, _ = x.shape\n\n        cls_tokens = repeat(self.cls_token, '() n d -> b n d', b = b)\n        x = torch.cat((cls_tokens, x), dim=1)\n        x += self.pos_embedding[:, :(n + 1)]\n        x = self.dropout(x)\n\n        x = self.transformer(x, mask)\n\n        x = self.to_cls_token(x[:, 0])\n        \n        return self.mlp_head(x)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:53:57.903401Z","iopub.execute_input":"2023-04-09T19:53:57.903767Z","iopub.status.idle":"2023-04-09T19:54:09.381131Z","shell.execute_reply.started":"2023-04-09T19:53:57.903733Z","shell.execute_reply":"2023-04-09T19:54:09.379795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_vit = ViT(\n    image_size = 7,\n    patch_size = 1,\n    num_classes = 16,\n    dim = 1024,\n    depth = 6,\n    heads = 16,\n    mlp_dim = 2048,\n    channels = 24,\n    dropout = 0.1,\n    emb_dropout = 0.1\n).cuda()","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:54:09.385414Z","iopub.execute_input":"2023-04-09T19:54:09.387519Z","iopub.status.idle":"2023-04-09T19:54:09.933187Z","shell.execute_reply.started":"2023-04-09T19:54:09.387483Z","shell.execute_reply":"2023-04-09T19:54:09.932019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from IPython import display\ndef train(net, model_vit, train_iter, valida_iter, loss, optimizer, device, epochs, early_stopping=True,\n          early_num=20):\n    loss_list = [100]\n    early_epoch = 0\n\n    net = net.to(device)\n    print(\"training on \", device)\n    start = time.time()\n    train_loss_list = []\n    valida_loss_list = []\n    train_acc_list = []\n    valida_acc_list = []\n    for epoch in range(epochs):\n        train_acc_sum, n = 0.0, 0\n        time_epoch = time.time()\n        lr_adjust = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, 15, eta_min=0.0, last_epoch=-1)\n        for X, y in train_iter:\n            \n            batch_count, train_l_sum = 0, 0\n            #X = X.permute(0, 3, 1, 2)\n            X = X.to(device)\n            y = y.to(device)\n            y_hat = net(X)\n            y_hat = model_vit(y_hat)\n            # print('y_hat:', y_hat.shape)\n            # print('y:', y.shape)\n            l = loss(y_hat, y.long())\n\n            optimizer.zero_grad()\n            l.backward()\n            optimizer.step()\n            train_l_sum += l.cpu().item()\n            train_acc_sum += (y_hat.argmax(dim=1) == y).sum().cpu().item()\n            n += y.shape[0]\n            batch_count += 1\n        lr_adjust.step(epoch)\n        valida_acc, valida_loss = evaluate_accuracy(valida_iter, net, model_vit, loss, device)\n        loss_list.append(valida_loss)\n\n        \n        train_loss_list.append(train_l_sum) # / batch_count)\n        train_acc_list.append(train_acc_sum / n)\n        valida_loss_list.append(valida_loss)\n        valida_acc_list.append(valida_acc)\n\n        print('epoch %d, train loss %.6f, train acc %.3f, valida loss %.6f, valida acc %.3f, time %.1f sec'\n                % (epoch + 1, train_l_sum / batch_count, train_acc_sum / n, valida_loss, valida_acc, time.time() - time_epoch))\n\n        PATH = \"./net_DBA.pt\"\n        # if loss_list[-1] <= 0.01 and valida_acc >= 0.95:\n        #     torch.save(net.state_dict(), PATH)\n        #     break\n\n        if early_stopping and loss_list[-2] < loss_list[-1]:  # < 0.05) and (loss_list[-1] <= 0.05):\n            if early_epoch == 0: # and valida_acc > 0.9:\n                torch.save(net.state_dict(), PATH)\n            early_epoch += 1\n            loss_list[-1] = loss_list[-2]\n            if early_epoch == early_num:\n                net.load_state_dict(torch.load(PATH))\n                break\n        else:\n            early_epoch = 0\n\n    \n    set_figsize()\n    plt.figure(figsize=(8, 8.5))\n    train_accuracy =   plt.subplot(221)\n    train_accuracy.set_title('train_accuracy')\n    plt.plot(np.linspace(1, epoch, len(train_acc_list)), train_acc_list, color='green')\n    plt.xlabel('epoch')\n    plt.ylabel('train_accuracy')\n    \n    test_accuracy =   plt.subplot(222)\n    test_accuracy.set_title('valida_accuracy')\n    plt.plot(np.linspace(1, epoch, len(valida_acc_list)), valida_acc_list, color='deepskyblue')\n    plt.xlabel('epoch')\n    plt.ylabel('test_accuracy')\n\n    loss_sum =   plt.subplot(223)\n    loss_sum.set_title('train_loss')\n    plt.plot(np.linspace(1, epoch, len(train_loss_list)), train_loss_list, color='red')\n    plt.xlabel('epoch')\n    plt.ylabel('train loss')\n    # ls_plot = np.array(ls_plot)\n\n    test_loss =   plt.subplot(224)\n    test_loss.set_title('valida_loss')\n    plt.plot(np.linspace(1, epoch, len(valida_loss_list)), valida_loss_list, color='gold')\n    plt.xlabel('epoch')\n    plt.ylabel('valida loss')\n    # ls_plot = np.array(ls_plot)\n\n    plt.show()\n    print('epoch %d, loss %.4f, train acc %.3f, time %.1f sec'\n            % (epoch + 1, train_l_sum / batch_count, train_acc_sum / n, time.time() - start))\n","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:54:12.592151Z","iopub.execute_input":"2023-04-09T19:54:12.592531Z","iopub.status.idle":"2023-04-09T19:54:12.612707Z","shell.execute_reply.started":"2023-04-09T19:54:12.592495Z","shell.execute_reply":"2023-04-09T19:54:12.611409Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def sampling(proportion, ground_truth, bg=False):\n    train = {}\n    test = {}\n    labels_loc = {}\n    m = max(ground_truth)+1\n    for i in range(m):\n        a = 0 if bg == True else 1\n        indexes = [j for j, x in enumerate(ground_truth.ravel().tolist()) if x == i + a]\n        np.random.shuffle(indexes)\n        labels_loc[i] = indexes\n        if proportion != 1:\n            nb_val = max(int((1 - proportion) * len(indexes)), 3)\n        else:\n            nb_val = 0\n        train[i] = indexes[:nb_val]\n        test[i] = indexes[nb_val:]\n        \n        gc.collect()\n    train_indexes = []\n    test_indexes = []\n    for i in range(m):\n        train_indexes += train[i]\n        test_indexes += test[i]\n    np.random.shuffle(train_indexes)\n    np.random.shuffle(test_indexes)\n    return train_indexes, test_indexes\n","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:55:04.734210Z","iopub.execute_input":"2023-04-09T19:55:04.734609Z","iopub.status.idle":"2023-04-09T19:55:04.744237Z","shell.execute_reply.started":"2023-04-09T19:55:04.734572Z","shell.execute_reply":"2023-04-09T19:55:04.742907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install torch-optimizer\nimport torch_optimizer as optim2","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:55:10.035197Z","iopub.execute_input":"2023-04-09T19:55:10.035574Z","iopub.status.idle":"2023-04-09T19:55:20.494688Z","shell.execute_reply.started":"2023-04-09T19:55:10.035538Z","shell.execute_reply":"2023-04-09T19:55:20.493379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import ctypes\nlibc = ctypes.CDLL(\"libc.so.6\") # clearing cache \nlibc.malloc_trim(0)","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:55:20.497181Z","iopub.execute_input":"2023-04-09T19:55:20.497562Z","iopub.status.idle":"2023-04-09T19:55:20.517690Z","shell.execute_reply.started":"2023-04-09T19:55:20.497519Z","shell.execute_reply":"2023-04-09T19:55:20.516744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import time\nimport collections\nfrom torch import optim\n\nres = None\n\nfor index_iter in range(ITER):\n    print('iter:', index_iter)\n    #define the model\n    #net = pResNet(32, 48, CLASSES_NUM, BAND, 2, 16, bottleneck=True)\n    #net = resnet20(num_classes=CLASSES_NUM)\n    net = SSRN_network(BAND, CLASSES_NUM)\n\n    #optimizer = optim2.DiffGrad(net.parameters(), lr=lr, amsgrad=False) #, weight_decay=0.0001)\n    optimizer = torch.optim.Adam(list(net.parameters())+list(model_vit.parameters()), lr= 1e-3, betas=(0.9, 0.999),\n    eps=1e-8,\n    weight_decay=0)\n    time_1 = int(time.time())\n    np.random.seed(seeds[index_iter])\n    train_indices, test_indices = sampling(VALIDATION_SPLIT, gt)\n    _, total_indices = sampling(1, gt)\n    _, total_indicesbg = sampling(1, gt, bg=True)\n    TOTAL_SIZEBG=21025\n\n    TRAIN_SIZE = len(train_indices)\n    print('Train size: ', TRAIN_SIZE)\n    TEST_SIZE = TOTAL_SIZE - TRAIN_SIZE\n    print('Test size: ', TEST_SIZE)\n    VAL_SIZE = int(TRAIN_SIZE)\n    print('Validation size: ', VAL_SIZE)\n\n    print('-----Selecting Small Pieces from the Original Cube Data-----')\n    train_iter, valida_iter, test_iter, all_iter, all_iter_bg = \\\n    generate_iter(TRAIN_SIZE, train_indices, TEST_SIZE, test_indices, TOTAL_SIZE, total_indices, TOTAL_SIZEBG, total_indicesbg, VAL_SIZE, whole_data, PATCH_LENGTH, padded_data, INPUT_DIMENSION, 16, gt) #batchsize in 1\n\n    tic1 = time.time()\n    train(net, model_vit, train_iter, valida_iter, loss, optimizer, device, epochs=100)\n    toc1 = time.time()\n\n    pred_test = []\n    tic2 = time.time()\n    with torch.no_grad():\n        for X, y in test_iter:\n            #X = X.permute(0, 3, 1, 2)\n            X = X.to(device)\n            net.eval()\n            model_vit.eval()\n            y_hat = net(X)\n            y_hat = model_vit(y_hat)\n            pred_test.extend(np.array(y_hat.cpu().argmax(axis=1)))\n    toc2 = time.time()\n    collections.Counter(pred_test)\n    gt_test = gt[test_indices] - 1\n\n\n    overall_acc = metrics.accuracy_score(pred_test, gt_test[:-VAL_SIZE])\n    confusion_matrix = metrics.confusion_matrix(pred_test, gt_test[:-VAL_SIZE])\n    each_acc, average_acc = aa_and_each_accuracy(confusion_matrix)\n    kappa = metrics.cohen_kappa_score(pred_test, gt_test[:-VAL_SIZE])\n\n    torch.save(net.state_dict(), \"./content/\" + str(round(overall_acc, 3)) + '.pt')\n    KAPPA.append(kappa)\n    OA.append(overall_acc)\n    AA.append(average_acc)\n    TRAINING_TIME.append(toc1 - tic1)\n    TESTING_TIME.append(toc2 - tic2)\n    ELEMENT_ACC[index_iter, :] = each_acc\n\n","metadata":{"execution":{"iopub.status.busy":"2023-04-09T19:55:35.826529Z","iopub.execute_input":"2023-04-09T19:55:35.826926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"--------\" + \" Training Finished-----------\")\nrecord_output(OA, AA, KAPPA, ELEMENT_ACC, TRAINING_TIME, TESTING_TIME,confusion_matrix,\n                     './content/'  + str(img_rows) + '_' + Dataset + 'split：' + str(VALIDATION_SPLIT) + 'lr：' + str(lr) + '.txt')\n\n\n#generate_png(all_iter, net, model_vit, gt_hsi, Dataset, device, total_indices)\ngenerate_png(all_iter_bg, net, model_vit, gt_hsi, Dataset, device, total_indicesbg)","metadata":{},"execution_count":null,"outputs":[]}]}