{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":25563,"databundleVersionId":2094376,"sourceType":"competition"},{"sourceId":2032065,"sourceType":"datasetVersion","datasetId":1216613}],"dockerImageVersionId":30699,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport os\n\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import DataLoader,Dataset\nfrom torchvision import transforms\nfrom torchvision import models as models\nimport torchvision\nfrom torch.utils.data import random_split  \nimport torch\nimport torch.nn as nn\n\nfrom matplotlib import pyplot as plt\nfrom PIL import Image\n\nimport albumentations as A # 图像增强\nfrom albumentations.pytorch import ToTensorV2 # 图像增强\n\nfrom tqdm.notebook import tqdm\n\nprint(\"import done\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-17T09:12:41.367340Z","iopub.execute_input":"2024-05-17T09:12:41.367733Z","iopub.status.idle":"2024-05-17T09:12:49.738858Z","shell.execute_reply.started":"2024-05-17T09:12:41.367695Z","shell.execute_reply":"2024-05-17T09:12:49.737901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'  \nprint(\"device:\",device)\n\nclass_threshold = 0.4\n\nnum_epoch = 20\nbatch_size = 64\nlr = 0.0001\nimg_size = 224  ","metadata":{"execution":{"iopub.status.busy":"2024-05-17T09:12:49.740865Z","iopub.execute_input":"2024-05-17T09:12:49.741517Z","iopub.status.idle":"2024-05-17T09:12:49.772870Z","shell.execute_reply.started":"2024-05-17T09:12:49.741481Z","shell.execute_reply":"2024-05-17T09:12:49.771866Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 1. 处理数据集\nresize_img_file = '../input/resized-plant2021/img_sz_256'\ntest_img_dir =  '../input/plant-pathology-2021-fgvc8/test_images/'\ntrain_origin = pd.read_csv('../input/plant-pathology-2021-fgvc8/train.csv')\n\ntrain_origin.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-17T09:12:54.238377Z","iopub.execute_input":"2024-05-17T09:12:54.238998Z","iopub.status.idle":"2024-05-17T09:12:54.286078Z","shell.execute_reply.started":"2024-05-17T09:12:54.238969Z","shell.execute_reply":"2024-05-17T09:12:54.285208Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 2. 多标签label使用one-hot初始化\nlabels_list = [\"healthy\", \"scab\", \"rust\", \"frog_eye_leaf_spot\", \"powdery_mildew\", \"complex\"]\ntrain_df = train_origin[['image']].copy()\nfor label in labels_list:\n    train_df[label] = 0\n\nfor label in labels_list:\n    train_df.loc[train_origin['labels'].str.contains(label), label] = 1\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-17T09:12:56.432975Z","iopub.execute_input":"2024-05-17T09:12:56.433627Z","iopub.status.idle":"2024-05-17T09:12:56.516264Z","shell.execute_reply.started":"2024-05-17T09:12:56.433583Z","shell.execute_reply":"2024-05-17T09:12:56.515279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PathologyPlantsDataset(Dataset): # 加载图像\n    def __init__(self, image_ids, targets, path, mode, transform=None):\n        self.image_ids = image_ids\n        self.targets = targets\n        self.root_dir = path\n        self.mode = mode\n        self.transform = transform\n    \n    def __len__(self):\n        return len(self.image_ids)\n    \n    def __getitem__(self, idx):\n        # 读取图像\n        image_path = os.path.join(self.root_dir, self.image_ids.iloc[idx])\n        image = Image.open(image_path)\n        img = np.array(image)\n        # 处理图像\n        if self.transform:\n            image = self.transform(image=img)['image']\n        \n        if self.mode == 'test':\n            target = None\n        else:\n            target = torch.tensor(self.targets[idx], dtype=torch.float32) \n        \n        return (image, target)","metadata":{"execution":{"iopub.status.busy":"2024-05-17T09:12:58.638834Z","iopub.execute_input":"2024-05-17T09:12:58.639218Z","iopub.status.idle":"2024-05-17T09:12:58.647667Z","shell.execute_reply.started":"2024-05-17T09:12:58.639191Z","shell.execute_reply":"2024-05-17T09:12:58.646378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transform = A.Compose([\n    A.Rotate(\n        always_apply=False, \n        p=0.1, \n        limit=(-68, 178), \n        interpolation=1, \n        border_mode=0, \n        value=(0, 0, 0), \n        mask_value=None\n    ),\n    A.RandomShadow(\n        num_shadows_lower=1, \n        num_shadows_upper=1, \n        shadow_dimension=3, \n        shadow_roi=(0, 0.6, 1, 1), \n        p=0.4\n    ),\n    A.ShiftScaleRotate(\n        shift_limit=0.05, \n        scale_limit=0.05, \n        rotate_limit=15, \n        p=0.6\n    ),\n    A.RandomFog(\n        fog_coef_lower=0.2, \n        fog_coef_upper=0.2, \n        alpha_coef=0.2, \n        p=0.3\n    ),\n    A.RGBShift(\n        r_shift_limit=15, \n        g_shift_limit=15, \n        b_shift_limit=15, \n        p=0.3\n    ),\n    A.RandomBrightnessContrast(\n        p=0.3\n    ),\n    A.GaussNoise(\n        var_limit=(50, 70),  \n        always_apply=False, \n        p=0.3\n    ),\n    A.Resize(\n        height=img_size,\n        width=img_size,\n    ),\n    A.CoarseDropout(\n        max_holes=5, \n        max_height=5, \n        max_width=5, \n        min_holes=3, \n        min_height=5, \n        min_width=5,\n        always_apply=False, \n        p=0.2\n    ),\n    A.Normalize(\n        mean=(0.485, 0.456, 0.406), \n        std=(0.229, 0.224, 0.225)\n    ),\n    ToTensorV2(),\n])\nval_transform = A.Compose([\n    A.Resize(\n        height=img_size,\n        width=img_size,\n    ),\n    A.Normalize(\n        mean=(0.485, 0.456, 0.406), \n        std=(0.229, 0.224, 0.225)\n    ),\n    ToTensorV2(),\n])\n\n\nX_train, X_valid, y_train, y_valid = train_test_split(\n    train_df['image'], \n    train_df[labels_list].values,  \n    test_size=0.2, \n    random_state=42\n)\nprint(\"训练数据集长度：\",len(X_train))\nprint(\"测试数据集长度：\",len(X_valid))\n\nfor i in range(2):  # 打印前2个样本 检查\n    print(f\"Image path: {X_train.iloc[i]}, Labels: {y_train[i]}\") ","metadata":{"execution":{"iopub.status.busy":"2024-05-17T09:28:57.866168Z","iopub.execute_input":"2024-05-17T09:28:57.866538Z","iopub.status.idle":"2024-05-17T09:28:57.886348Z","shell.execute_reply.started":"2024-05-17T09:28:57.866509Z","shell.execute_reply":"2024-05-17T09:28:57.885354Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_set = PathologyPlantsDataset(X_train, y_train, path=resize_img_file ,mode='train',transform=train_transform)\nval_set = PathologyPlantsDataset(X_valid, y_valid, path=resize_img_file, mode='valid', transform=val_transform)\n\n# 超参设置\ntrain_loader = DataLoader(train_set, batch_size=batch_size, shuffle=True)\nvalid_loader = DataLoader(val_set, batch_size=batch_size, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2024-05-17T09:28:59.747696Z","iopub.execute_input":"2024-05-17T09:28:59.748054Z","iopub.status.idle":"2024-05-17T09:28:59.754104Z","shell.execute_reply.started":"2024-05-17T09:28:59.748027Z","shell.execute_reply":"2024-05-17T09:28:59.753222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def model_resnet():\n    # 加载预训练的 ResNet50 模型\n    model = models.resnet50(weights='ResNet50_Weights.DEFAULT')\n    \n    # 解冻所有层的参数（如果需要微调整个模型，可以注释掉下面这段代码）\n    for param in model.parameters():\n        param.requires_grad = False\n\n    # 修改模型的最后一层\n    model.fc = nn.Sequential(\n        nn.Linear(model.fc.in_features,512),\n        nn.ReLU(),  # ReLU 激活函数\n        nn.BatchNorm1d(512),  # 批标准化层\n        nn.Dropout(0.5),\n        nn.Linear(512, 6),  # 最终的全连接层，将特征维度映射到类别数\n        nn.Sigmoid()  # Sigmoid 激活函数，用于多标签分类\n    )\n    \n    return model\n","metadata":{"execution":{"iopub.status.busy":"2024-05-17T09:29:01.701006Z","iopub.execute_input":"2024-05-17T09:29:01.701973Z","iopub.status.idle":"2024-05-17T09:29:01.708422Z","shell.execute_reply.started":"2024-05-17T09:29:01.701937Z","shell.execute_reply":"2024-05-17T09:29:01.707216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"my_resnet = model_resnet().to(device)\n# loss_fn = torch.nn.MultiLabelSoftMarginLoss()\nloss_fn = nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(my_resnet.parameters(), lr=lr)","metadata":{"execution":{"iopub.status.busy":"2024-05-17T09:29:03.512552Z","iopub.execute_input":"2024-05-17T09:29:03.513403Z","iopub.status.idle":"2024-05-17T09:29:04.102257Z","shell.execute_reply.started":"2024-05-17T09:29:03.513367Z","shell.execute_reply":"2024-05-17T09:29:04.101236Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import f1_score, accuracy_score\nfrom tqdm import tqdm\n\ndef to_numpy(tensor):\n    return tensor.detach().cpu().numpy()\n\ndef get_metrics(y_pred_proba,y_test,threshold=0.5):\n    y_pred = np.where(y_pred_proba > threshold, 1, 0)\n    y1 = y_pred.round().astype(np.float32)\n    y2 = y_test.round().astype(np.float32)\n    \n    \n    y_pred = (y_pred_proba > threshold).astype(int)  # 直接转换为整数  \n    f1 = f1_score(y_test, y_pred, average='samples', zero_division=1)\n    acc = accuracy_score(y_test,y_pred, normalize=True)\n\n    return acc, f1 \n\ndef train_or_valid(dataloader, model, device, loss_fn, optimizer=None, is_train=True):\n    loss_val = 0\n    accuracy = 0\n    f1score = 0\n    num_batches = len(dataloader)\n    \n    if is_train:\n        model.train()\n    else:\n        model.eval()\n    \n    with torch.set_grad_enabled(is_train):\n        stream = tqdm(dataloader)\n        for batch, (X, y) in enumerate(stream, start=1):\n            X, y = X.to(device), y.to(device)\n            \n            pred_prob = model(X)\n            loss = loss_fn(pred_prob, y)\n            \n            if is_train:\n                optimizer.zero_grad()\n                loss.backward()\n                optimizer.step()\n            \n            loss_val += loss.item()\n            acc, f1 = get_metrics(to_numpy(pred_prob), to_numpy(y))\n            \n            accuracy += acc\n            f1score += f1\n            \n            desc = f'Epoch {epoch:3d}/{num_epochs} - {\"train\" if is_train else \"valid\"}_Loss: {loss_val/batch:.4f}, ' + \\\n                   f'{\"train\" if is_train else \"valid\"}_Acc: {accuracy/batch:.4f}, {\"train\" if is_train else \"valid\"}_F1: {f1score/batch:.4f}'\n            stream.set_description(desc)\n    \n    return loss_val / num_batches, accuracy / num_batches, f1score / num_batches\n\ndef train(dataloader, model, device, loss_fn, optimizer, train_loss, train_acc, train_f1, epoch, num_epochs):\n    loss, acc, f1 = train_or_valid(dataloader, model, device, loss_fn, optimizer, is_train=True)\n    train_loss.append(loss)\n    train_acc.append(acc)\n    train_f1.append(f1)\n\ndef valid(dataloader, model, device, loss_fn, valid_loss, valid_acc, valid_f1, epoch, num_epochs):\n    loss, acc, f1 = train_or_valid(dataloader, model, device, loss_fn, optimizer=None, is_train=False)\n    valid_loss.append(loss)\n    valid_acc.append(acc)\n    valid_f1.append(f1)\n    return f1\n","metadata":{"execution":{"iopub.status.busy":"2024-05-17T09:29:05.608535Z","iopub.execute_input":"2024-05-17T09:29:05.609167Z","iopub.status.idle":"2024-05-17T09:29:05.625892Z","shell.execute_reply.started":"2024-05-17T09:29:05.609135Z","shell.execute_reply":"2024-05-17T09:29:05.624965Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\n# 训练过程\nnum_epochs = 20\ntrain_loss, train_acc, train_f1 = [], [], []\nvalid_loss, valid_acc, valid_f1 = [], [], []\n\nbest_f1 = 0\nfor epoch in range(1, num_epochs + 1):\n    train(train_loader, my_resnet, device, loss_fn, optimizer, train_loss, train_acc, train_f1, epoch, num_epochs)\n    vaild_f1 = valid(valid_loader, my_resnet, device, loss_fn, valid_loss, valid_acc, valid_f1, epoch, num_epochs)\n    if vaild_f1 > best_f1:\n        torch.save(my_resnet.state_dict(),'Best_model.pth')\n        best_f1 = vaild_f1","metadata":{"execution":{"iopub.status.busy":"2024-05-17T09:29:07.852771Z","iopub.execute_input":"2024-05-17T09:29:07.853394Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from matplotlib.ticker import MaxNLocator \n\ndef plot_result(train,valid,mode,file_name):\n    epochs = range(1, len(train) + 1)\n    fig, ax = plt.subplots(figsize=(8, 5)) \n    if mode == 'loss':\n        ax.plot(epochs, train, label='Training loss', marker='o')  \n        ax.plot(epochs, valid, label='Validation loss', marker='o')  \n        ax.legend(frameon=False, fontsize=14)  \n        ax.get_xaxis().set_major_locator(MaxNLocator(integer=True))  \n        ax.set_title('Loss', fontsize=18)  \n        ax.set_xlabel('Epoch', fontsize=14)  \n        ax.set_ylabel('Loss', fontsize=14)  \n        plt.savefig(file_name + '.png')\n        plt.close(fig)\n    elif mode == 'acc':\n        ax.plot(epochs, train, label='Training Accuracy', marker='o')  \n        ax.plot(epochs, valid, label='Validation accuracy', marker='o')  \n        ax.legend(frameon=False, fontsize=14)  \n        ax.get_xaxis().set_major_locator(MaxNLocator(integer=True))  \n        ax.set_title('Accuracy', fontsize=18)  \n        ax.set_xlabel('Epoch', fontsize=14)  \n        ax.set_ylabel('Accuracy', fontsize=14)  \n        plt.savefig(file_name + '.png')\n        plt.close(fig)\n    elif mode =='f1':\n        ax.plot(epochs, train, label='Training F1-Score', marker='o')  \n        ax.plot(epochs, valid, label='Validation F1-Score', marker='o')  \n        ax.legend(frameon=False, fontsize=14)  \n        ax.get_xaxis().set_major_locator(MaxNLocator(integer=True))  \n        ax.set_title('F1-Score', fontsize=18)  \n        ax.set_xlabel('Epoch', fontsize=14)  \n        ax.set_ylabel('F1-Score', fontsize=14)  \n        plt.savefig(file_name + '.png')\n        plt.close(fig)\n                ","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_result(train_loss,valid_loss,'loss','res_loss')    \nplot_result(train_acc,valid_acc,'acc','res_acc')  \nplot_result(train_f1,valid_f1,'f1','res_f1')  ","metadata":{},"execution_count":null,"outputs":[]}]}