{"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":"! pip install timm","metadata":{"execution":{"iopub.status.busy":"2022-09-23T14:26:42.954538Z","iopub.execute_input":"2022-09-23T14:26:42.955163Z","iopub.status.idle":"2022-09-23T14:26:54.970432Z","shell.execute_reply.started":"2022-09-23T14:26:42.955048Z","shell.execute_reply":"2022-09-23T14:26:54.969264Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🥅GPU信息","metadata":{"papermill":{"duration":0.007762,"end_time":"2022-08-19T03:18:40.613842","exception":false,"start_time":"2022-08-19T03:18:40.60608","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!nvidia-smi","metadata":{"papermill":{"duration":1.135846,"end_time":"2022-08-19T03:18:41.755854","exception":false,"start_time":"2022-08-19T03:18:40.620008","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T14:26:54.974565Z","iopub.execute_input":"2022-09-23T14:26:54.974863Z","iopub.status.idle":"2022-09-23T14:26:55.967231Z","shell.execute_reply.started":"2022-09-23T14:26:54.974831Z","shell.execute_reply":"2022-09-23T14:26:55.966102Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🚀导入工具包","metadata":{"papermill":{"duration":0.006162,"end_time":"2022-08-19T03:18:41.768517","exception":false,"start_time":"2022-08-19T03:18:41.762355","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import os\nimport gc\nimport cv2\nimport copy\nimport time\nimport timm\nimport torch\nimport random\nimport string\nimport joblib\nimport tifffile\nimport numpy as np \nimport pandas as pd \nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\n\nfrom torch import nn\nfrom torchvision import models\nfrom tqdm.notebook import tqdm\nfrom torchvision import transforms\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.utils.tensorboard import SummaryWriter\nfrom sklearn.model_selection import train_test_split\n\nbar = '=='\nos.makedirs('./log', exist_ok=True)\nwriter = SummaryWriter('./log')\ndevice = torch.device('cuda') if torch.cuda.is_available() else 'cpu'\nprint(bar*20)\nprint(f'PyTorch Version :{torch.__version__}')\nprint(f'Device :{device}')","metadata":{"papermill":{"duration":7.926297,"end_time":"2022-08-19T03:18:49.700954","exception":false,"start_time":"2022-08-19T03:18:41.774657","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T14:26:55.96939Z","iopub.execute_input":"2022-09-23T14:26:55.969801Z","iopub.status.idle":"2022-09-23T14:27:03.565221Z","shell.execute_reply.started":"2022-09-23T14:26:55.969756Z","shell.execute_reply":"2022-09-23T14:27:03.563521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ⛳路径PATH","metadata":{"papermill":{"duration":0.00568,"end_time":"2022-08-19T03:18:49.713619","exception":false,"start_time":"2022-08-19T03:18:49.707939","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# CSV\ntrain_csv_path = '../input/mayo-competition-dataset/Mayo_Competition_V2/train.csv'\ntest_csv_path = '../input/mayo-clinic-strip-ai/test.csv'\n\n# Image\ntrain_image = '../input/mayo-competition-dataset/Mayo_Competition_V2/train_224/'\ntest_image  = '../input/jpg-images-strip-ai/test/'\n","metadata":{"papermill":{"duration":0.0141,"end_time":"2022-08-19T03:18:49.73337","exception":false,"start_time":"2022-08-19T03:18:49.71927","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T14:27:03.568226Z","iopub.execute_input":"2022-09-23T14:27:03.568826Z","iopub.status.idle":"2022-09-23T14:27:03.574471Z","shell.execute_reply.started":"2022-09-23T14:27:03.568791Z","shell.execute_reply":"2022-09-23T14:27:03.573245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🕹一些开关","metadata":{"papermill":{"duration":0.00577,"end_time":"2022-08-19T03:18:49.745182","exception":false,"start_time":"2022-08-19T03:18:49.739412","status":"completed"},"tags":[]}},{"cell_type":"code","source":"debug = False # 控制train_csv的显示数目\ngenerate_new = False # 是否转换图像","metadata":{"papermill":{"duration":0.014738,"end_time":"2022-08-19T03:18:49.765626","exception":false,"start_time":"2022-08-19T03:18:49.750888","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T14:27:03.577467Z","iopub.execute_input":"2022-09-23T14:27:03.577727Z","iopub.status.idle":"2022-09-23T14:27:03.583659Z","shell.execute_reply.started":"2022-09-23T14:27:03.577703Z","shell.execute_reply":"2022-09-23T14:27:03.582614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🎿读取CSV文件的数据","metadata":{"papermill":{"duration":0.005476,"end_time":"2022-08-19T03:18:49.776722","exception":false,"start_time":"2022-08-19T03:18:49.771246","status":"completed"},"tags":[]}},{"cell_type":"code","source":"train_csv_debug = pd.read_csv(train_csv_path).head(10 if debug else 1000)\ntrain_csv = pd.read_csv(train_csv_path)\ntest_csv = pd.read_csv(test_csv_path)","metadata":{"papermill":{"duration":0.03218,"end_time":"2022-08-19T03:18:49.814592","exception":false,"start_time":"2022-08-19T03:18:49.782412","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T14:27:03.585521Z","iopub.execute_input":"2022-09-23T14:27:03.585932Z","iopub.status.idle":"2022-09-23T14:27:03.622962Z","shell.execute_reply.started":"2022-09-23T14:27:03.585898Z","shell.execute_reply":"2022-09-23T14:27:03.622109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_csv)","metadata":{"papermill":{"duration":0.017377,"end_time":"2022-08-19T03:18:49.838422","exception":false,"start_time":"2022-08-19T03:18:49.821045","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T14:27:03.624381Z","iopub.execute_input":"2022-09-23T14:27:03.624691Z","iopub.status.idle":"2022-09-23T14:27:03.633255Z","shell.execute_reply.started":"2022-09-23T14:27:03.62466Z","shell.execute_reply":"2022-09-23T14:27:03.632051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### 查看数据Test_csv文件中","metadata":{"papermill":{"duration":0.006195,"end_time":"2022-08-19T03:18:49.85043","exception":false,"start_time":"2022-08-19T03:18:49.844235","status":"completed"},"tags":[]}},{"cell_type":"code","source":"test_csv","metadata":{"papermill":{"duration":0.025512,"end_time":"2022-08-19T03:18:49.881741","exception":false,"start_time":"2022-08-19T03:18:49.856229","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T14:27:03.635174Z","iopub.execute_input":"2022-09-23T14:27:03.63599Z","iopub.status.idle":"2022-09-23T14:27:03.654311Z","shell.execute_reply.started":"2022-09-23T14:27:03.635955Z","shell.execute_reply":"2022-09-23T14:27:03.653484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 👌转换图像","metadata":{"papermill":{"duration":0.005685,"end_time":"2022-08-19T03:18:49.893809","exception":false,"start_time":"2022-08-19T03:18:49.888124","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if(generate_new):\n    # 分别创建俩个文件夹用来存放已转换的图像数据\n    os.makedirs(\"./train/\")\n    os.makedirs(\"./test/\")\n    \n    # 转换测试用的图像\n    for i in tqdm(range(test_df.shape[0])): \n        \n        img_id = test_csv.iloc[i].image_id # 图像id\n        \n        img = cv2.resize(tifffile.imread(test_image + img_id + \".tif\"), (512, 512))# 转换出图像\n        cv2.imwrite(f\"./test/{img_id}.jpg\", img)\n        del img\n        gc.collect()\n    \n    # 转换训练的图像\n    for i in tqdm(range(train_df.shape[0])):\n        \n        img_id = train_csv.iloc[i].image_id # 图像id\n        img = cv2.resize(tifffile.imread(train_image + img_id + \".tif\"), (512, 512))# 转换出图像\n        cv2.imwrite(f\"./train/{img_id}.jpg\", img)\n        del img # 从内存中删除防止占用过多的内存\n        gc.collect()","metadata":{"papermill":{"duration":0.017004,"end_time":"2022-08-19T03:18:49.917288","exception":false,"start_time":"2022-08-19T03:18:49.900284","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T14:27:03.65534Z","iopub.execute_input":"2022-09-23T14:27:03.65634Z","iopub.status.idle":"2022-09-23T14:27:03.665477Z","shell.execute_reply.started":"2022-09-23T14:27:03.656289Z","shell.execute_reply":"2022-09-23T14:27:03.664566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🐱‍🏍图像的增强与处理","metadata":{}},{"cell_type":"markdown","source":"### 绘制一些原始图像 ","metadata":{"papermill":{"duration":0.006342,"end_time":"2022-08-19T03:18:49.962578","exception":false,"start_time":"2022-08-19T03:18:49.956236","status":"completed"},"tags":[]}},{"cell_type":"code","source":"images = []\nlabels = []\nims = os.listdir('../input/mayo-competition-dataset/Mayo_Competition_V2/train_224/')\n\nfor i in ims :\n    image = cv2.imread('../input/mayo-competition-dataset/Mayo_Competition_V2/train_224/'+i)\n    images.append(image)\n    \nfor i in range(len(train_csv)) :\n    label = {\"CE\" : 0, \"LAA\": 1}[train_csv.iloc[i].label]\n    labels.append(label)","metadata":{"papermill":{"duration":9.364643,"end_time":"2022-08-19T03:18:59.333231","exception":false,"start_time":"2022-08-19T03:18:49.968588","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T14:27:03.670134Z","iopub.execute_input":"2022-09-23T14:27:03.670417Z","iopub.status.idle":"2022-09-23T14:27:10.740796Z","shell.execute_reply.started":"2022-09-23T14:27:03.670392Z","shell.execute_reply":"2022-09-23T14:27:10.739801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"* 大动脉粥样硬化性卒中（LAA）\n* 心源性脑栓塞（CE）","metadata":{"papermill":{"duration":0.006011,"end_time":"2022-08-19T03:18:59.345737","exception":false,"start_time":"2022-08-19T03:18:59.339726","status":"completed"},"tags":[]}},{"cell_type":"code","source":"debug = False\nshow_num = 5\nnums_images = 25#len(images)\nlabel_name = ['CE', 'LAA']\n\nplt.figure(figsize=(10, 10))\nfor i in range(nums_images):\n    ax = plt.subplot(show_num, show_num, i + 1)\n    if debug:\n        plt.imshow(images[i])\n        plt.axis(\"on\")\n        \n    else :\n        name = label_name[labels[i]]\n        ax.set_title(name, fontproperties='SimHei', fontsize=20)\n        plt.imshow(images[i])\n        plt.axis(\"off\")","metadata":{"papermill":{"duration":1.727738,"end_time":"2022-08-19T03:19:01.079424","exception":false,"start_time":"2022-08-19T03:18:59.351686","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T14:27:10.742448Z","iopub.execute_input":"2022-09-23T14:27:10.742838Z","iopub.status.idle":"2022-09-23T14:27:11.941346Z","shell.execute_reply.started":"2022-09-23T14:27:10.742797Z","shell.execute_reply":"2022-09-23T14:27:11.940394Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 数据增强","metadata":{}},{"cell_type":"code","source":"# 重新组合颜色通道\ndef change_channel(img):\n    b = cv2.split(img)[0]\n    g = cv2.split(img)[1]\n    r = cv2.split(img)[2]\n    brg = cv2.merge([b, r, g]) # 可以自己改变组合顺序\n    return brg\n\n# 添加高斯噪声\ndef noise(image, mean = 0, var = 0.01):\n    '''\n        添加高斯噪声\n        mean : 均值\n        var : 方差，方差越大越模糊\n    '''\n    image = np.array(image/255, dtype=float)\n    noise = np.random.normal(mean, var ** 0.5, image.shape)\n    out = image + noise\n    if out.min() < 0:\n        low_clip = -1.\n    else:\n        low_clip = 0.\n    out = np.clip(out, low_clip, 1.0)\n    out = np.uint8(out*255)\n    return out","metadata":{"execution":{"iopub.status.busy":"2022-09-23T14:27:11.942555Z","iopub.execute_input":"2022-09-23T14:27:11.944592Z","iopub.status.idle":"2022-09-23T14:27:11.954152Z","shell.execute_reply.started":"2022-09-23T14:27:11.94456Z","shell.execute_reply":"2022-09-23T14:27:11.952989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"im = []\nfor i in ims :\n    image = cv2.imread('../input/mayo-competition-dataset/Mayo_Competition_V2/train_224/'+i)\n    image = change_channel(image)\n    im.append(image)\n    ","metadata":{"execution":{"iopub.status.busy":"2022-09-23T14:27:11.956013Z","iopub.execute_input":"2022-09-23T14:27:11.956377Z","iopub.status.idle":"2022-09-23T14:27:13.63035Z","shell.execute_reply.started":"2022-09-23T14:27:11.956342Z","shell.execute_reply":"2022-09-23T14:27:13.629336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"* **增强后图像噪声+颜色通道重构**","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(10, 10))\nfor i in range(nums_images):\n    ax = plt.subplot(show_num, show_num, i + 1)\n    if debug:\n        \n        plt.imshow(im[i])\n        plt.axis(\"on\")\n        \n    else :\n        name = label_name[labels[i]]\n        ax.set_title(name, fontproperties='SimHei', fontsize=20)\n        plt.imshow(im[i])\n        plt.axis(\"off\")\ndel im\ndel images","metadata":{"execution":{"iopub.status.busy":"2022-09-23T14:27:13.631744Z","iopub.execute_input":"2022-09-23T14:27:13.632127Z","iopub.status.idle":"2022-09-23T14:27:14.586171Z","shell.execute_reply.started":"2022-09-23T14:27:13.632091Z","shell.execute_reply":"2022-09-23T14:27:14.585222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🎓制作图像数据集","metadata":{"papermill":{"duration":0.006066,"end_time":"2022-08-19T03:18:49.929274","exception":false,"start_time":"2022-08-19T03:18:49.923208","status":"completed"},"tags":[]}},{"cell_type":"code","source":"len(train_csv)","metadata":{"execution":{"iopub.status.busy":"2022-09-23T14:27:14.587488Z","iopub.execute_input":"2022-09-23T14:27:14.588234Z","iopub.status.idle":"2022-09-23T14:27:14.595603Z","shell.execute_reply.started":"2022-09-23T14:27:14.588197Z","shell.execute_reply":"2022-09-23T14:27:14.594512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ImgDataset(Dataset):\n    def __init__(self, csv_file):\n        self.csv = csv_file \n        self.train = 'label' in csv_file.columns\n            \n    def __len__(self):\n        return len(self.csv)\n    \n    def __getitem__(self, index):\n        # 选择路径\n        if(generate_new):\n            paths = [\"./test/\", \"./train/\"]\n        else:\n            paths = [\"../input/mayo-competition-dataset/Mayo_Competition_V2/test_224/\",\n                     \"../input/mayo-competition-dataset/Mayo_Competition_V2/train_224/\"]\n        # 读取图像\n        image_names = os.listdir(paths[self.train])\n        \n        image = cv2.imread(paths[self.train] + image_names[index]) # 拼接出图像的路径并读取\n        # 图像增强 \n        image = change_channel(image) # 颜色通道重构\n        # 转置图像\n        if len(image.shape) == 5:\n            image = image.squeeze().transpose(1, 2, 0)\n        # 调整图像的大小\n        image = cv2.resize(image, (512, 512)).transpose(2, 0, 1) # (3,512,512)\n        label = None # 平常的标签一般为none也就是无标签\n        \n        if(self.train): # 如果存在标签在CSV文件中就应用标签\n            label = {\"CE\" : 0, \"LAA\": 1}[self.csv.iloc[index].label]\n            \n        return image, label","metadata":{"papermill":{"duration":0.021974,"end_time":"2022-08-19T03:19:01.112719","exception":false,"start_time":"2022-08-19T03:19:01.090745","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T14:27:14.597728Z","iopub.execute_input":"2022-09-23T14:27:14.598083Z","iopub.status.idle":"2022-09-23T14:27:14.610766Z","shell.execute_reply.started":"2022-09-23T14:27:14.598048Z","shell.execute_reply":"2022-09-23T14:27:14.609422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **训练函数**","metadata":{"papermill":{"duration":0.011181,"end_time":"2022-08-19T03:19:01.134534","exception":false,"start_time":"2022-08-19T03:19:01.123353","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def train_model(model, dataloaders_config, criterion, optimizer, EPOCHS, Loss_List):\n    best_acc = 0.0 \n    lr_list = []\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer=optimizer, T_0=5)\n    print(f'初始化的学习率 :',optimizer.defaults['lr'])\n\n    for epoch in range(EPOCHS):\n        model.to(device)\n        ######################################################\n        for phase in ['train', 'val']: # 选择验证还是训练\n            ##################\n            if phase == 'train': # 如果是训练\n                model.train() # 将模型的状态设定为训练模式\n            else:\n                model.eval() # 将模型设为推理模式\n            ##################  \n            epoch_loss = 0.0 # 迭代的损失计数器\n            epoch_acc = 0 # 迭代的准确率计数器\n            step = 0 # step计数器\n            ##################\n            dataloader = dataloaders_config[phase] # 选择数据加载器（其中包含是训练数据集的加载器还是验证数据集的加载器）\n            for Images, Labels in tqdm(dataloader, leave=False):\n                images = Images.to(device).float() # 将图像数据放到合适的设备\n                labels = Labels.to(device).long() # 将标签数据放到合适的设备\n                optimizer.zero_grad() # 归零优化器的梯度/优化器进行初始化\n                \n                with torch.set_grad_enabled(phase == 'train'): # 设置是否求导开关如果为tran则将对模型的参数进行求导否则为val时则不求导\n                    model_output = model(images) # 将图像数据喂入模型\n                    loss = criterion(model_output, labels) # 使用损失函数机算模型的输出预测和对应标签的损失\n                    _, preds = torch.max(model_output, 1) # 在模型的输出的预测值中筛选出数值最大的\n\n                    ##########\n                    if phase == 'train': # 如果为训练模式则进行损失的反向传播和执行优化器对模型的参数进行更新\n                        loss.backward() # 反向传播\n                        optimizer.step() # 执行优化器\n                        lr_list.append(optimizer.param_groups[0]['lr'])\n                        scheduler.step() # 使用学习率衰减器\n                        \n                        writer.add_scalar('Lr :', optimizer.param_groups[0]['lr'], global_step=step)\n                    ##########\n                    epoch_loss += loss.item() * len(model_output)\n                    epoch_acc += torch.sum(preds == labels.data)\n                step += 1 \n                   \n                 \n            ####################################\n            data_size = len(dataloader.dataset)\n            epoch_loss = epoch_loss / data_size\n            epoch_acc = epoch_acc.double() / data_size\n            print(f'Epoch {epoch + 1}/{EPOCHS} | {phase:^5} | Loss: {epoch_loss:.4f} | Acc: {epoch_acc:.4f}')\n            \n            Loss_List.append(epoch_loss)\n            writer.add_scalar('Epochs Loss', epoch_loss, global_step=epoch)\n            writer.add_scalar('Epochs Acc', epoch_acc, global_step=epoch)\n        ######################################################\n        if epoch_acc > best_acc:\n            #save_model = torch.save(model.state_dict(),'./model.pth') # 保存模型\n            \n            saved_model = torch.jit.script(model)\n            saved_model.save('saved_model.pt')\n            best_acc = epoch_acc","metadata":{"papermill":{"duration":0.027659,"end_time":"2022-08-19T03:19:01.173066","exception":false,"start_time":"2022-08-19T03:19:01.145407","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T14:27:14.612584Z","iopub.execute_input":"2022-09-23T14:27:14.612928Z","iopub.status.idle":"2022-09-23T14:27:14.627924Z","shell.execute_reply.started":"2022-09-23T14:27:14.612894Z","shell.execute_reply":"2022-09-23T14:27:14.627309Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 与训练有关的配置","metadata":{"papermill":{"duration":0.010527,"end_time":"2022-08-19T03:19:01.195253","exception":false,"start_time":"2022-08-19T03:19:01.184726","status":"completed"},"tags":[]}},{"cell_type":"code","source":"batch_size = 1\ntrain, val = train_test_split(train_csv, test_size=0.1, random_state=42, stratify = train_csv.label)","metadata":{"papermill":{"duration":0.025726,"end_time":"2022-08-19T03:19:01.231879","exception":false,"start_time":"2022-08-19T03:19:01.206153","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T14:27:14.629033Z","iopub.execute_input":"2022-09-23T14:27:14.631368Z","iopub.status.idle":"2022-09-23T14:27:14.646308Z","shell.execute_reply.started":"2022-09-23T14:27:14.631313Z","shell.execute_reply":"2022-09-23T14:27:14.645231Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train)","metadata":{"execution":{"iopub.status.busy":"2022-09-23T14:27:14.647635Z","iopub.execute_input":"2022-09-23T14:27:14.648247Z","iopub.status.idle":"2022-09-23T14:27:14.65492Z","shell.execute_reply.started":"2022-09-23T14:27:14.648211Z","shell.execute_reply":"2022-09-23T14:27:14.654016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🎈数据加载器","metadata":{"papermill":{"duration":0.010916,"end_time":"2022-08-19T03:19:01.253565","exception":false,"start_time":"2022-08-19T03:19:01.242649","status":"completed"},"tags":[]}},{"cell_type":"code","source":"train_loader = DataLoader(\n    ImgDataset(train_csv), \n    batch_size=batch_size, \n    shuffle=True, \n    num_workers=1\n)\n\n\nval_loader = DataLoader(\n    ImgDataset(val), \n    batch_size=batch_size, \n    shuffle=False, \n    num_workers=1\n)\n\n\ndataloaders_config = {\"train\": train_loader, \"val\": val_loader} # 分别配置数据加载器","metadata":{"papermill":{"duration":0.018566,"end_time":"2022-08-19T03:19:01.283153","exception":false,"start_time":"2022-08-19T03:19:01.264587","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T14:27:14.656256Z","iopub.execute_input":"2022-09-23T14:27:14.656841Z","iopub.status.idle":"2022-09-23T14:27:14.663948Z","shell.execute_reply.started":"2022-09-23T14:27:14.656805Z","shell.execute_reply":"2022-09-23T14:27:14.662993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🎲人工神经网络模型","metadata":{"papermill":{"duration":0.010574,"end_time":"2022-08-19T03:19:01.304204","exception":false,"start_time":"2022-08-19T03:19:01.29363","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class BelndModel(nn.Module) :\n    def __init__(self) :\n        super(BelndModel, self).__init__()\n        \n        self.M1 = torch.hub.load('NVIDIA/DeepLearningExamples:torchhub', 'nvidia_efficientnet_b4', pretrained=True)\n        self.M1.classifier.fc = nn.Linear(1792, 2, bias=True)\n        \n        self.M2 = models.regnet_x_16gf(pretrained=True)\n        self.M2.fc = nn.Linear(2048, 2, bias=True)\n        \n        self.M3 = timm.create_model(\"convnext_xlarge_384_in22ft1k\", pretrained=True, num_classes=2)\n        \n        \n    def forward(self, x) :\n        \n        #print(x.dtype)\n        x = x.float()\n        #print(x.shape)\n        \n        x1 = self.M1(x)\n        x2 = self.M2(x)\n        \n        x3 = transforms.functional.resize(x, size=[384, 384]) / 255.0\n        x3 = transforms.functional.normalize(x3,\n                                             mean=[0.48145466, 0.4578275, 0.40821073], \n                                             std=[0.26862954, 0.26130258, 0.27577711])\n        \n        x3 = self.M3(x3)\n        \n        y = (x1*1.0) + (x2* 1.0) \n        y = y + (x3*0.9)\n        \n        return y\n        \n        ","metadata":{"execution":{"iopub.status.busy":"2022-09-23T14:27:14.665444Z","iopub.execute_input":"2022-09-23T14:27:14.665791Z","iopub.status.idle":"2022-09-23T14:27:14.675443Z","shell.execute_reply.started":"2022-09-23T14:27:14.665758Z","shell.execute_reply":"2022-09-23T14:27:14.674541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"c = torch.hub.load('NVIDIA/DeepLearningExamples:torchhub', 'nvidia_efficientnet_b4', pretrained=True)\nc.classifier","metadata":{"execution":{"iopub.status.busy":"2022-09-23T14:27:14.676747Z","iopub.execute_input":"2022-09-23T14:27:14.677404Z","iopub.status.idle":"2022-09-23T14:27:28.282908Z","shell.execute_reply.started":"2022-09-23T14:27:14.677361Z","shell.execute_reply":"2022-09-23T14:27:28.281966Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"c.classifier.fc = nn.Linear(1792, 2, bias=True)\nc.classifier","metadata":{"execution":{"iopub.status.busy":"2022-09-23T14:27:28.284236Z","iopub.execute_input":"2022-09-23T14:27:28.285156Z","iopub.status.idle":"2022-09-23T14:27:28.294596Z","shell.execute_reply.started":"2022-09-23T14:27:28.285113Z","shell.execute_reply":"2022-09-23T14:27:28.292638Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#model = torch.hub.load('NVIDIA/DeepLearningExamples:torchhub', 'nvidia_efficientnet_b4', pretrained=True)\n#model = models.regnet_x_16gf()\n#model.fc = nn.Linear(2048, 2, bias=True)\n\nmodel = BelndModel()\n\n\ncriterion = nn.CrossEntropyLoss() # 选择损失函数\noptimizer_0 = torch.optim.AdamW(model.parameters(), lr=1e-4) # 优化器_0\noptimizer_1 = torch.optim.AdamW(model.parameters(), lr=1e-5) # 优化器_1","metadata":{"papermill":{"duration":8.552441,"end_time":"2022-08-19T03:19:09.867231","exception":false,"start_time":"2022-08-19T03:19:01.31479","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T14:27:28.296366Z","iopub.execute_input":"2022-09-23T14:27:28.297237Z","iopub.status.idle":"2022-09-23T14:28:49.708071Z","shell.execute_reply.started":"2022-09-23T14:27:28.297199Z","shell.execute_reply":"2022-09-23T14:28:49.707102Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🥇开始训练（分开训练也就是执行两个train()并从中挑选合适的模型）","metadata":{"papermill":{"duration":0.011024,"end_time":"2022-08-19T03:19:09.890169","exception":false,"start_time":"2022-08-19T03:19:09.879145","status":"completed"},"tags":[]}},{"cell_type":"code","source":"epochs_loss_list = []\ntrain_model(model,\n            dataloaders_config,\n            criterion,\n            optimizer_0,\n            3,\n            Loss_List=epochs_loss_list)\n","metadata":{"papermill":{"duration":226.112489,"end_time":"2022-08-19T03:22:56.013466","exception":false,"start_time":"2022-08-19T03:19:09.900977","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-23T14:28:49.709657Z","iopub.execute_input":"2022-09-23T14:28:49.71005Z","iopub.status.idle":"2022-09-23T14:42:13.899052Z","shell.execute_reply.started":"2022-09-23T14:28:49.710012Z","shell.execute_reply":"2022-09-23T14:42:13.897452Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.title('Train Loss')\nplt.plot(epochs_loss_list)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-09-23T14:42:13.917663Z","iopub.execute_input":"2022-09-23T14:42:13.920764Z","iopub.status.idle":"2022-09-23T14:42:14.352817Z","shell.execute_reply.started":"2022-09-23T14:42:13.920706Z","shell.execute_reply":"2022-09-23T14:42:14.351771Z"},"trusted":true},"execution_count":null,"outputs":[]}]}