{"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":"none","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"}],"dockerImageVersionId":30664,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport torch\nimport torchvision.transforms as transforms\nfrom torchvision.models import inception_v3\nfrom torch.utils.data import Dataset, IterableDataset, DataLoader\nimport torch.nn.functional as F\nfrom torch import nn\nfrom tqdm import tqdm\nimport math\nfrom torchvision import transforms\n","metadata":{"execution":{"iopub.status.busy":"2024-03-05T02:16:41.179332Z","iopub.execute_input":"2024-03-05T02:16:41.180060Z","iopub.status.idle":"2024-03-05T02:16:44.852757Z","shell.execute_reply.started":"2024-03-05T02:16:41.180020Z","shell.execute_reply":"2024-03-05T02:16:44.851749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"以下这几块是试验代码","metadata":{}},{"cell_type":"code","source":"\"\"\"import numpy as np\nimport matplotlib.pyplot as plt\n\nparquet = pd.read_parquet(\n    '/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/1000913311.parquet'\n)\n\nX = parquet['Fp1'] - parquet['F7']\n# 生成示例序列X\nfs = 200  # 采样率\nprint(len(X))\nt = np.linspace(0, len(X) // fs, fs, endpoint=False)  # 生成时间序列\nprint(X.shape)\n# 计算谱图\nres = plt.specgram(X, NFFT=256, Fs=fs, noverlap=128, cmap='viridis')\nprint(res[0])\nprint(res[0].shape, res[1].shape, res[2].shape)\n\n# 绘制谱图\nplt.xlabel('time')\nplt.ylabel('freq(Hz)')\nplt.title('putu')\nplt.colorbar(label='qiangdu')\nplt.show()\"\"\"","metadata":{"execution":{"iopub.status.busy":"2024-03-05T02:16:44.854806Z","iopub.execute_input":"2024-03-05T02:16:44.855247Z","iopub.status.idle":"2024-03-05T02:16:44.863311Z","shell.execute_reply.started":"2024-03-05T02:16:44.855216Z","shell.execute_reply":"2024-03-05T02:16:44.861985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\nimport torch\n\n\n# 假设 tensor_8_channel 是你的8通道张量，形状为 [8, 100, 100]\ntensor_8_channel = torch.randn((1, 2, 2))\n\nprint(tensor_8_channel)\n\n# 1. 裁剪或者调整大小为3通道，299乘299\ntensor_cropped_resized = transforms.functional.resize(tensor_8_channel, (3, 3))\n\nprint(tensor_cropped_resized)\n\n\"\"\"","metadata":{"execution":{"iopub.status.busy":"2024-03-05T02:16:44.864492Z","iopub.execute_input":"2024-03-05T02:16:44.864845Z","iopub.status.idle":"2024-03-05T02:16:44.887302Z","shell.execute_reply.started":"2024-03-05T02:16:44.864811Z","shell.execute_reply":"2024-03-05T02:16:44.886408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"试验代码结束","metadata":{}},{"cell_type":"markdown","source":"# 数据读取","metadata":{}},{"cell_type":"code","source":"train_file = '../input/hms-harmful-brain-activity-classification/train.csv'\nparquet_file = '../input/hms-harmful-brain-activity-classification/train_eegs/'\ntrain_data_mode = 0\ndrop_rate = 0.4\n\ntarget_p = [\n    'seizure_vote', \n    'lpd_vote', \n    'gpd_vote', \n    'lrda_vote', \n    'grda_vote', \n    'other_vote'\n]\n\nfeatures = [\n    'Fp1','T3','C3','O1','Fp2','C4','T4','O2'\n]","metadata":{"execution":{"iopub.status.busy":"2024-03-05T02:16:44.889821Z","iopub.execute_input":"2024-03-05T02:16:44.890657Z","iopub.status.idle":"2024-03-05T02:16:44.900507Z","shell.execute_reply.started":"2024-03-05T02:16:44.890629Z","shell.execute_reply":"2024-03-05T02:16:44.899766Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"如果eeg_id不合并，而是每个eeg_sub_id都生成一个样本呢？\n\n如果每个不同的eeg_id只生成一个样本，train_data_mode将设为1，如果eeg_id不合并，train_data_mode将设为0","metadata":{}},{"cell_type":"code","source":"train = pd.read_csv(train_file)\ntrain","metadata":{"execution":{"iopub.status.busy":"2024-03-05T02:16:44.901472Z","iopub.execute_input":"2024-03-05T02:16:44.901762Z","iopub.status.idle":"2024-03-05T02:16:45.094297Z","shell.execute_reply.started":"2024-03-05T02:16:44.901738Z","shell.execute_reply":"2024-03-05T02:16:45.093325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train['total_evaluators'] = train[target_p].sum(axis=1)\n\nif train_data_mode == 1:\n    # 将数据框按照 eeg_id 进行分组，并对每个分组进行合并操作\n    train_grouped_eeg = train.groupby('eeg_id').agg({\n        'eeg_sub_id': 'count',  # 计算每个 eeg_id 对应的行数\n        'eeg_label_offset_seconds': 'first',\n        'spectrogram_id': 'first',\n        'spectrogram_sub_id': 'first',\n        'spectrogram_label_offset_seconds': 'first',\n        'label_id': 'first',\n        'patient_id': 'first',\n        'expert_consensus': 'first',\n        'seizure_vote': 'sum',\n        'lpd_vote': 'sum',\n        'gpd_vote': 'sum',\n        'lrda_vote': 'sum',\n        'grda_vote': 'sum',\n        'other_vote': 'sum',\n        'total_evaluators': 'sum'\n    }).reset_index()\nelse:\n    train_grouped_eeg = train\n    \ntrain_grouped_eeg['consensus'] = train_grouped_eeg[target_p].max(axis=1)\n\n# 计算每一列的比例\ntrain_grouped_eeg[target_p] = train_grouped_eeg[target_p].div(\n    train_grouped_eeg['total_evaluators'], \n    axis=0\n)\n    \ntrain_grouped_eeg","metadata":{"execution":{"iopub.status.busy":"2024-03-05T02:16:45.095634Z","iopub.execute_input":"2024-03-05T02:16:45.096260Z","iopub.status.idle":"2024-03-05T02:16:45.174096Z","shell.execute_reply.started":"2024-03-05T02:16:45.096222Z","shell.execute_reply":"2024-03-05T02:16:45.172992Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"按drop_rate去掉不需要的行","metadata":{}},{"cell_type":"code","source":"train_grouped_eeg['row_agreement'] = train_grouped_eeg['consensus'] / train_grouped_eeg['total_evaluators']\ntrain_grouped_eeg = train_grouped_eeg[train_grouped_eeg['row_agreement'] >= drop_rate]\ntrain_grouped_eeg","metadata":{"execution":{"iopub.status.busy":"2024-03-05T02:16:45.175387Z","iopub.execute_input":"2024-03-05T02:16:45.175698Z","iopub.status.idle":"2024-03-05T02:16:45.215508Z","shell.execute_reply.started":"2024-03-05T02:16:45.175671Z","shell.execute_reply":"2024-03-05T02:16:45.214663Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HMSDataset(IterableDataset):\n    def __init__(self, data, path):\n        super(HMSDataset, self).__init__()\n        self.data = data\n        self.path = path\n    \n    def __iter__(self):\n        for i in range(len(self.data)):\n            x = self.load_sample(i)\n            y = torch.tensor(\n                self.data.iloc[i][target_p].to_numpy().astype(np.float32), \n                dtype=torch.float32\n            )\n            yield x, y\n    \n    def load_sample(self, index):\n        eeg_id = self.data.iloc[index]['eeg_id']\n        parquet = pd.read_parquet(\n            self.path + str(eeg_id) + '.parquet', \n            columns = features\n        )\n        #print(\"load \" + str(eeg_id) + '.parquet')\n        \"\"\"# TODO:新公式(4, 100, 100)\n        data = []\n        data[0] = parquet[''] - parquet['']\n        data[1] = parquet[''] - parquet['']\n        data[2] = parquet[''] - parquet['']\n        data[3] = parquet[''] - parquet['']\n        data[4] = parquet[''] - parquet['']\n        data[5] = parquet[''] - parquet['']\n        data[6] = parquet[''] - parquet['']\n        data[7] = parquet[''] - parquet['']\"\"\"\n        \n        start = int(50 * self.data.iloc[index]['eeg_label_offset_seconds'])\n        parquet_data = parquet[start: start+ 10000].to_numpy()\n        parquet_data = np.nan_to_num(parquet_data, nan=0.0)\n        parquet_data = torch.tensor(parquet_data, dtype=torch.float32).permute(1, 0)\n        parquet_data = parquet_data.contiguous().reshape(len(features), 100, 100)\n        return parquet_data","metadata":{"execution":{"iopub.status.busy":"2024-03-05T02:17:15.637715Z","iopub.execute_input":"2024-03-05T02:17:15.638517Z","iopub.status.idle":"2024-03-05T02:17:15.648083Z","shell.execute_reply.started":"2024-03-05T02:17:15.638484Z","shell.execute_reply":"2024-03-05T02:17:15.647026Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hmsdataset = HMSDataset(train_grouped_eeg, parquet_file)\ndataloader = DataLoader(hmsdataset, batch_size=64)","metadata":{"execution":{"iopub.status.busy":"2024-03-05T02:19:39.553152Z","iopub.execute_input":"2024-03-05T02:19:39.553910Z","iopub.status.idle":"2024-03-05T02:19:39.558991Z","shell.execute_reply.started":"2024-03-05T02:19:39.553875Z","shell.execute_reply":"2024-03-05T02:19:39.557977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target_p\ntrain_grouped_eeg.iloc[0][target_p]","metadata":{"execution":{"iopub.status.busy":"2024-03-05T02:19:41.561336Z","iopub.execute_input":"2024-03-05T02:19:41.562400Z","iopub.status.idle":"2024-03-05T02:19:41.571724Z","shell.execute_reply.started":"2024-03-05T02:19:41.562354Z","shell.execute_reply":"2024-03-05T02:19:41.570751Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for x, y in dataloader:\n    print(x.shape, y.shape)\n    break","metadata":{"execution":{"iopub.status.busy":"2024-03-05T02:19:43.190403Z","iopub.execute_input":"2024-03-05T02:19:43.190788Z","iopub.status.idle":"2024-03-05T02:19:43.603495Z","shell.execute_reply.started":"2024-03-05T02:19:43.190756Z","shell.execute_reply":"2024-03-05T02:19:43.602527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MInceptionV3(nn.Module):\n    def __init__(self, input_shape, num_class):\n        super(MInceptionV3, self).__init__()\n        self.upconv = nn.ConvTranspose2d(input_shape[0], input_shape[0], 4, 3)\n        self.input_conv = nn.Conv2d(input_shape[0], 3, 3)\n        self.inception_v3 = inception_v3(init_weights=True)\n        self.inception_v3.fc = nn.Linear(self.inception_v3.fc.in_features, num_class)\n        \n    def forward(self, x):\n        y = self.input_conv(self.upconv(x))\n        outputs, _ = self.inception_v3(y)\n        return outputs","metadata":{"execution":{"iopub.status.busy":"2024-03-05T02:16:45.343732Z","iopub.execute_input":"2024-03-05T02:16:45.344043Z","iopub.status.idle":"2024-03-05T02:16:45.351212Z","shell.execute_reply.started":"2024-03-05T02:16:45.344015Z","shell.execute_reply":"2024-03-05T02:16:45.350152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomCrossEntropyLoss(nn.Module):\n    def __init__(self):\n        super(CustomCrossEntropyLoss, self).__init__()\n        \n    def forward(self, y_hat, y):\n        y_hat = F.log_softmax(y_hat, dim=-1)\n        return -torch.sum(y * y_hat) / y.shape[0]","metadata":{"execution":{"iopub.status.busy":"2024-03-05T02:16:45.352619Z","iopub.execute_input":"2024-03-05T02:16:45.353176Z","iopub.status.idle":"2024-03-05T02:16:45.364168Z","shell.execute_reply.started":"2024-03-05T02:16:45.353139Z","shell.execute_reply":"2024-03-05T02:16:45.363209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Trainer:\n    def __init__(self, epoch, model, optimizer, criterion, device):\n        self.epoch = epoch\n        self.model = model\n        self.optimizer = optimizer\n        self.criterion = criterion\n        self.device = device\n        \n    def train(self, dataloader):\n        train_loss = []\n        for i in range(self.epoch):\n            loss = self.train_epoch(dataloader)\n            print(f'{i}th epoch, loss is {loss}')\n            train_loss.append(loss)\n            \n    def train_epoch(self, dataloader):\n        total_loss = 0.0\n        for batch in tqdm(dataloader, desc='Processing', unit='batch'):\n            x, y = batch\n            x = x.to(device)\n            y = y.to(device)\n            loss = self.train_step(x, y)\n            if math.isnan(loss):\n                print(loss)\n            total_loss += loss\n        return total_loss / len(dataloader)\n            \n    def train_step(self, inputs, targets):\n        self.model.train()\n        self.optimizer.zero_grad()\n        outputs = self.model(inputs)\n        loss = self.criterion(outputs, targets)\n        loss.backward()\n        optimizer.step()\n        return loss.item()\n    ","metadata":{"execution":{"iopub.status.busy":"2024-03-05T02:16:45.365276Z","iopub.execute_input":"2024-03-05T02:16:45.365580Z","iopub.status.idle":"2024-03-05T02:16:45.375624Z","shell.execute_reply.started":"2024-03-05T02:16:45.365554Z","shell.execute_reply":"2024-03-05T02:16:45.374683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"input_shape = (len(features), 100, 100)\ndevice = 'cuda'\nmodel = MInceptionV3(input_shape, len(target_p)).to('cuda')\noptimizer = torch.optim.Adam(model.parameters(), lr=0.001)\ncriterion = CustomCrossEntropyLoss()\ntrainer = Trainer(10, model, optimizer, criterion, device)","metadata":{"execution":{"iopub.status.busy":"2024-03-05T02:16:45.376659Z","iopub.execute_input":"2024-03-05T02:16:45.376918Z","iopub.status.idle":"2024-03-05T02:16:46.096266Z","shell.execute_reply.started":"2024-03-05T02:16:45.376895Z","shell.execute_reply":"2024-03-05T02:16:46.095414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.train(dataloader)","metadata":{"execution":{"iopub.status.busy":"2024-03-05T02:19:48.153389Z","iopub.execute_input":"2024-03-05T02:19:48.154335Z","iopub.status.idle":"2024-03-05T02:27:19.672719Z","shell.execute_reply.started":"2024-03-05T02:19:48.154298Z","shell.execute_reply":"2024-03-05T02:27:19.671396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), 'inceptionv3_24_3_1.pth')","metadata":{"execution":{"iopub.status.busy":"2024-03-05T02:17:00.855536Z","iopub.status.idle":"2024-03-05T02:17:00.856025Z","shell.execute_reply.started":"2024-03-05T02:17:00.855766Z","shell.execute_reply":"2024-03-05T02:17:00.855785Z"},"trusted":true},"execution_count":null,"outputs":[]}]}