{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":13836,"databundleVersionId":1718836,"sourceType":"competition"},{"sourceId":7431789,"sourceType":"datasetVersion","datasetId":4324811}],"dockerImageVersionId":30635,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport random\nimport cv2\nimport os\nfrom tqdm import tqdm\nfrom glob import glob\nfrom sklearn.model_selection import GroupKFold, StratifiedKFold, StratifiedGroupKFold\nimport warnings\nimport time\nfrom matplotlib import pyplot as plt\n\nimport torch\nimport torch.nn as nn\nfrom torchvision import models, transforms\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn.functional as F\nfrom torch.cuda.amp import autocast, GradScaler\n\nwarnings.filterwarnings('ignore')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-01-19T06:46:51.914386Z","iopub.execute_input":"2024-01-19T06:46:51.915189Z","iopub.status.idle":"2024-01-19T06:46:51.923326Z","shell.execute_reply.started":"2024-01-19T06:46:51.915151Z","shell.execute_reply":"2024-01-19T06:46:51.922053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 简介\n> 在本次竞赛中，我们介绍了在乌干达定期调查中收集的 21,367 张标签图像数据集。大部分图像都是由农民拍摄的自家菜园照片众包而来，并由国家作物资源研究所（NaCRRI）的专家与坎帕拉马凯雷雷大学的人工智能实验室合作进行标注。这种格式最真实地反映了农民在现实生活中需要诊断的内容。\n\n> 您的任务是将每张木薯图像分为四个疾病类别或表示健康叶片的第五个类别。有了你们的帮助，农民们也许就能快速识别病株，从而有可能在病株造成不可挽回的损失之前挽救他们的庄稼。","metadata":{}},{"cell_type":"markdown","source":"# 分析数据","metadata":{}},{"cell_type":"code","source":"train_csv_path = '/kaggle/input/cassava-leaf-disease-classification/train.csv'\ndf = pd.read_csv(train_csv_path)\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:46:51.928299Z","iopub.execute_input":"2024-01-19T06:46:51.928721Z","iopub.status.idle":"2024-01-19T06:46:51.964539Z","shell.execute_reply.started":"2024-01-19T06:46:51.928674Z","shell.execute_reply":"2024-01-19T06:46:51.963544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_counts = df['label'].value_counts()\nlabel_counts","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:46:51.966444Z","iopub.execute_input":"2024-01-19T06:46:51.966814Z","iopub.status.idle":"2024-01-19T06:46:51.976139Z","shell.execute_reply.started":"2024-01-19T06:46:51.966783Z","shell.execute_reply":"2024-01-19T06:46:51.975032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 创建柱状图\nplt.bar(label_counts.index, label_counts.values)\n\n# 添加标签和标题\nplt.xlabel('Label')\nplt.ylabel('Count')\nplt.title('Label Distribution')\n\n# 显示图形\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:46:51.977528Z","iopub.execute_input":"2024-01-19T06:46:51.978301Z","iopub.status.idle":"2024-01-19T06:46:52.296575Z","shell.execute_reply.started":"2024-01-19T06:46:51.978257Z","shell.execute_reply":"2024-01-19T06:46:52.295389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 设置种子","metadata":{}},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n    print('种子设置成功！')\n\nseed_everything(42)","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:46:52.299879Z","iopub.execute_input":"2024-01-19T06:46:52.300327Z","iopub.status.idle":"2024-01-19T06:46:52.307896Z","shell.execute_reply.started":"2024-01-19T06:46:52.300295Z","shell.execute_reply":"2024-01-19T06:46:52.306780Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 全局参数和功能函数","metadata":{}},{"cell_type":"code","source":"class CFG:\n    seed          = 42\n    debug         = False  # 设置 debug=False 进行完整训练\n    exp_name      = 'leaf_disease_class'  # 实验名称\n    comment       = 'resnet32'  # 实验备注\n    model_name    = 'resnet32'  # 模型名称\n    backbone      = 'resnet'  # 模型骨干网络\n    train_bs      = 128  # 训练时的批量大小\n    valid_bs      = train_bs * 2  # 验证时的批量大小，通常是训练批量大小的两倍\n    img_size      = [224 * 2, 224 * 2]  # 输入图像的大小\n    epochs        = 50  # 训练的总轮次\n    lr            = 2e-3  # 初始学习率\n    scheduler     = 'CosineAnnealingLR'  # 学习率调度器的类型\n    min_lr        = 1e-6  # 学习率的最小值\n    T_max         = int(30000 / train_bs * epochs) + 50  # CosineAnnealingLR 调度器的周期\n    T_0           = 25  # CosineAnnealingWarmRestarts 调度器的初始周期\n    warmup_epochs = 0  # 学习率预热的轮次\n    wd            = 1e-6  # 权重衰减\n    n_accumulate  = max(1, 32 // train_bs)  # 累积梯度的步数，用于增大批量大小\n    n_fold        = 5  # 交叉验证的折数\n    num_classes   = 5  # 类别数\n    device        = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")  # 训练设备\n    \n    train_img_dir = '/kaggle/input/cassava-leaf-disease-classification/train_images'\n    test_img_dir  = '/kaggle/input/cassava-leaf-disease-classification/test_images'","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:46:52.309677Z","iopub.execute_input":"2024-01-19T06:46:52.310221Z","iopub.status.idle":"2024-01-19T06:46:52.319889Z","shell.execute_reply.started":"2024-01-19T06:46:52.310190Z","shell.execute_reply":"2024-01-19T06:46:52.318793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_img(img_path):\n    img = cv2.imread(img_path)\n    img = img[:,:,::-1].astype(np.float32)\n    mx = np.max(img)\n    if mx:\n        img = img/mx\n    return img\n\ndef show_img(img, label):\n    plt.figure()\n    plt.imshow(img)\n    plt.title(label)\n    plt.show()    ","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:46:52.321668Z","iopub.execute_input":"2024-01-19T06:46:52.322507Z","iopub.status.idle":"2024-01-19T06:46:52.331456Z","shell.execute_reply.started":"2024-01-19T06:46:52.322465Z","shell.execute_reply":"2024-01-19T06:46:52.330337Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test\ntmp_path = df.iloc[0, 0]\ntmp_label = df.iloc[0,1]\ntmp_path = os.path.join(CFG.train_img_dir, tmp_path)\nprint(tmp_path)\n\nimg = load_img(tmp_path)\nprint(img.shape)\nshow_img(img, tmp_label)","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:46:52.332817Z","iopub.execute_input":"2024-01-19T06:46:52.333227Z","iopub.status.idle":"2024-01-19T06:46:52.870937Z","shell.execute_reply.started":"2024-01-19T06:46:52.333196Z","shell.execute_reply":"2024-01-19T06:46:52.869897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 构建数据管道","metadata":{}},{"cell_type":"markdown","source":"> StratifiedKFold: StratifiedKFold是一种交叉验证方法，用于确保每个折叠中各类别样本的比例尽可能接近整体数据集中各类别样本的比例。这种方法特别适合于处理不平衡数据集。\n\n> split(np.arange(df.shape[0])): split方法是StratifiedKFold的核心，它实际上执行数据集的分割操作。np.arange(df.shape[0])生成一个从0到df.shape[0]（即数据集的行数）的整数序列，这个序列代表数据集中每一行的索引。df.label.values: 这是split方法的第二个参数，指定了数据集中的标签列。df.label.values提取了DataFrame df中名为label的列的值。这些标签用于确保在分层抽样过程中，每个折叠中各类别的样本比例反映整个数据集的比例。","metadata":{}},{"cell_type":"code","source":"folds = StratifiedKFold(n_splits=CFG.n_fold).split(np.arange(df.shape[0]),\n                                                      df.label.values)\nfor fold, (trn_idx, val_idx) in enumerate(folds):\n    df.loc[val_idx, 'fold'] = fold\n\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:46:52.872574Z","iopub.execute_input":"2024-01-19T06:46:52.872987Z","iopub.status.idle":"2024-01-19T06:46:52.898054Z","shell.execute_reply.started":"2024-01-19T06:46:52.872952Z","shell.execute_reply":"2024-01-19T06:46:52.897062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['fold'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:46:52.899484Z","iopub.execute_input":"2024-01-19T06:46:52.899952Z","iopub.status.idle":"2024-01-19T06:46:52.909669Z","shell.execute_reply.started":"2024-01-19T06:46:52.899898Z","shell.execute_reply":"2024-01-19T06:46:52.908571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display(df.groupby(['fold', 'label'])['image_id'].count())","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:46:52.914201Z","iopub.execute_input":"2024-01-19T06:46:52.914716Z","iopub.status.idle":"2024-01-19T06:46:52.933204Z","shell.execute_reply.started":"2024-01-19T06:46:52.914684Z","shell.execute_reply":"2024-01-19T06:46:52.932199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 定义 BuildDataset 类，继承自 torch.utils.data.Dataset\nclass BuildDataset(torch.utils.data.Dataset):\n    # 初始化方法，接收数据框 df、标签标志 label、和变换 transforms\n    def __init__(self, df, flag=True, transforms=None):\n        self.df         = df\n        self.flag      = flag\n        self.img_paths  = df['image_id'].tolist()\n        self.labels  = df['label'].tolist()\n        self.transforms = transforms\n\n    # 返回数据集的长度\n    def __len__(self):\n        return len(self.df)\n\n    # 根据给定的索引返回对应的样本\n    def __getitem__(self, index):\n        # 如果 label 为 True，表示需要加载标签\n        if self.flag:\n            # 获取图像路径\n            img_path = os.path.join(CFG.train_img_dir, self.img_paths[index])\n\n            # 调用 load_img 函数加载图像\n            img = load_img(img_path)\n        \n            # 获取标签\n            label = self.labels[index]\n            \n            # 如果提供了数据变换函数 transforms，则应用变换\n            if self.transforms:\n                data = self.transforms(image=img)\n                img  = data['image']\n            \n            # 将图像和掩码的通道维度调整为 (C, H, W) 的形状\n            img = np.transpose(img, (2, 0, 1))\n            \n            # 返回图像和掩码的 PyTorch 张量\n            return torch.tensor(img), torch.tensor(label)\n        else:\n            # 获取图像路径\n            img_path = os.path.join(CFG.test_img_dir, self.img_paths[index])\n\n            # 调用 load_img 函数加载图像\n            img = load_img(img_path)\n            \n            # 如果不需要加载标签，只加载图像\n            if self.transforms:\n                data = self.transforms(image=img)\n                img  = data['image']\n            \n            # 调整图像通道维度为 (C, H, W) 的形状\n            img = np.transpose(img, (2, 0, 1))\n            \n            # 返回图像的 PyTorch 张量\n            return torch.tensor(img)","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:46:52.934545Z","iopub.execute_input":"2024-01-19T06:46:52.934952Z","iopub.status.idle":"2024-01-19T06:46:52.948005Z","shell.execute_reply.started":"2024-01-19T06:46:52.934915Z","shell.execute_reply":"2024-01-19T06:46:52.946900Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import albumentations as A\n# 定义数据增强的转换管道\ndata_transforms = {\n    \"train\": A.Compose([\n        # 调整图像大小\n        A.Resize(*CFG.img_size, interpolation=cv2.INTER_NEAREST),\n        # 水平翻转\n        A.HorizontalFlip(p=0.5),\n        # 平移、缩放和旋转\n        A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.05, rotate_limit=10, p=0.5),\n        # 选择以下图像变换中的一种\n        A.OneOf([\n            # 网格扭曲\n            A.GridDistortion(num_steps=5, distort_limit=0.05, p=1.0),\n            # 弹性变换\n            A.ElasticTransform(alpha=1, sigma=50, alpha_affine=50, p=1.0)\n        ], p=0.25),\n        # 随机去除图像的部分区域\n        A.CoarseDropout(\n            max_holes=8, \n            max_height=CFG.img_size[0] // 20, \n            max_width=CFG.img_size[1] // 20,\n            min_holes=5, \n            fill_value=0, \n            mask_fill_value=0, \n            p=0.5\n        ),\n    ], p=1.0),\n    \n    \"valid\": A.Compose([\n        # 调整图像大小\n        A.Resize(*CFG.img_size, interpolation=cv2.INTER_NEAREST),\n    ], p=1.0)\n}","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:46:52.949397Z","iopub.execute_input":"2024-01-19T06:46:52.949725Z","iopub.status.idle":"2024-01-19T06:46:52.964794Z","shell.execute_reply.started":"2024-01-19T06:46:52.949698Z","shell.execute_reply":"2024-01-19T06:46:52.963746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 定义函数 prepare_loaders，用于准备训练和验证数据加载器\ndef prepare_loaders(fold):\n    # 从数据框中筛选出训练集和验证集\n    train_df = df.query(\"fold!=@fold\").reset_index(drop=True)\n    valid_df = df.query(\"fold==@fold\").reset_index(drop=True)\n    \n    # 创建训练集和验证集的数据集实例，应用相应的数据变换\n    train_dataset = BuildDataset(train_df, transforms=data_transforms['train'])\n    valid_dataset = BuildDataset(valid_df, transforms=data_transforms['valid'])\n\n    # 创建训练集和验证集的数据加载器\n    train_loader = DataLoader(train_dataset, batch_size=CFG.train_bs, \n                              num_workers=4, shuffle=True, pin_memory=True, drop_last=False)\n    valid_loader = DataLoader(valid_dataset, batch_size=CFG.valid_bs, \n                              num_workers=4, shuffle=False, pin_memory=True)\n    \n    # 返回训练集和验证集的数据加载器\n    return train_loader, valid_loader","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:46:52.966686Z","iopub.execute_input":"2024-01-19T06:46:52.967098Z","iopub.status.idle":"2024-01-19T06:46:52.980063Z","shell.execute_reply.started":"2024-01-19T06:46:52.967056Z","shell.execute_reply":"2024-01-19T06:46:52.979033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader, valid_loader = prepare_loaders(fold=0)\n\nimgs, labels = next(iter(train_loader))\nimgs.size(), labels.size()","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:46:52.981483Z","iopub.execute_input":"2024-01-19T06:46:52.982034Z","iopub.status.idle":"2024-01-19T06:47:19.037386Z","shell.execute_reply.started":"2024-01-19T06:46:52.981995Z","shell.execute_reply":"2024-01-19T06:47:19.035966Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_batch(imgs, labels, num_imgs=5, fig_width=15, fig_height=3):\n    \"\"\"\n    Plots a batch of images with their corresponding labels.\n    \"\"\"\n    plt.figure(figsize=(fig_width, fig_height))\n\n    # Ensure we do not try to plot more images than we have\n    num_imgs = min(num_imgs, len(imgs))\n\n    for idx in range(num_imgs):\n        plt.subplot(1, num_imgs, idx + 1)\n        img = imgs[idx].permute((1, 2, 0)).numpy() * 255\n        img = img.astype('uint8')\n        label = labels[idx]\n\n        plt.imshow(img)\n        plt.title(f\"Label: {label}\")\n        plt.axis('off')  # Turn off axis numbers and labels\n\n    plt.tight_layout()\n    plt.show()\n\n# Example usage\nplot_batch(imgs, labels, num_imgs=5, fig_width=15, fig_height=3)","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:19.041841Z","iopub.execute_input":"2024-01-19T06:47:19.042275Z","iopub.status.idle":"2024-01-19T06:47:20.050376Z","shell.execute_reply.started":"2024-01-19T06:47:19.042240Z","shell.execute_reply":"2024-01-19T06:47:20.049276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 在机器学习和数据处理的任务中，特别是处理大型数据集时，内存管理变得至关重要。\n# 垃圾回收可以帮助释放那些不再被引用的对象，从而释放内存，以便后续的计算能够更加高效\nimport gc\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:20.051983Z","iopub.execute_input":"2024-01-19T06:47:20.052331Z","iopub.status.idle":"2024-01-19T06:47:20.293199Z","shell.execute_reply.started":"2024-01-19T06:47:20.052300Z","shell.execute_reply":"2024-01-19T06:47:20.292008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 模型和损失函数","metadata":{}},{"cell_type":"code","source":"model = models.resnet34(pretrained=False) # 可以选择True，减少训练时间\nmodel.fc = nn.Linear(model.fc.in_features, CFG.num_classes)\nmodel","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-01-19T06:47:20.294857Z","iopub.execute_input":"2024-01-19T06:47:20.295300Z","iopub.status.idle":"2024-01-19T06:47:20.697046Z","shell.execute_reply.started":"2024-01-19T06:47:20.295261Z","shell.execute_reply":"2024-01-19T06:47:20.695913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model test\ntest_inputs = torch.randn(4,3,128,128)\nprint(model(test_inputs).shape)","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:20.698484Z","iopub.execute_input":"2024-01-19T06:47:20.698817Z","iopub.status.idle":"2024-01-19T06:47:20.910445Z","shell.execute_reply.started":"2024-01-19T06:47:20.698791Z","shell.execute_reply":"2024-01-19T06:47:20.909298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# nn.CrossEntropyLoss期望模型的输出是每个类的原始分数（即未经softmax处理的）。\n# 它同时进行log-softmax操作和计算交叉熵损失，这在数值上更稳定。\n# 标签应该是一个整数类型的Tensor，包含每个图像的类索引。\ncriterion = nn.CrossEntropyLoss()","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:20.911837Z","iopub.execute_input":"2024-01-19T06:47:20.912213Z","iopub.status.idle":"2024-01-19T06:47:20.917168Z","shell.execute_reply.started":"2024-01-19T06:47:20.912183Z","shell.execute_reply":"2024-01-19T06:47:20.916043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def calculate_accuracy(outputs, labels):\n    \"\"\"\n    计算准确率。\n\n    :param outputs: 模型的输出，大小为 (batch_size, num_classes)。\n    :param labels: 真实的标签，大小为 (batch_size)。\n    :return: 准确率。\n    \"\"\"\n    # 获取每个样本的预测类别。argmax 返回指定维度上最大值的索引\n    _, predicted = torch.max(outputs, 1)\n\n    # 计算预测正确的样本数量\n    correct = (predicted == labels).sum().item()\n\n    # 计算准确率\n    accuracy = correct / labels.size(0)\n    return accuracy\n\n# 示例使用\n# 假设 outputs 是模型的输出，labels 是真实标签\ntest_outputs = model(test_inputs)\ntest_labels = torch.ones(4)\naccuracy = calculate_accuracy(test_outputs, test_labels)\nprint(\"Accuracy:\", accuracy)","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:20.918701Z","iopub.execute_input":"2024-01-19T06:47:20.919029Z","iopub.status.idle":"2024-01-19T06:47:21.127439Z","shell.execute_reply.started":"2024-01-19T06:47:20.919001Z","shell.execute_reply":"2024-01-19T06:47:21.126273Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\n这个函数接受一个优化器 optimizer 作为参数，并根据配置文件中指定的 CFG.scheduler 的值选择相应的学习率调度器。\n支持的调度器有 CosineAnnealingLR、CosineAnnealingWarmRestarts、ReduceLROnPlateau、ExponentialLR。\n如果配置文件中的 CFG.scheduler 为 None，表示不使用学习率调度器，则返回 None。函数最终返回所选的学习率调度器。\n'''\ndef fetch_scheduler(optimizer):\n    # 根据配置选择不同的学习率调度器\n    if CFG.scheduler == 'CosineAnnealingLR':\n        scheduler = lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.T_max, \n                                                   eta_min=CFG.min_lr)\n    elif CFG.scheduler == 'CosineAnnealingWarmRestarts':\n        scheduler = lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=CFG.T_0, \n                                                             eta_min=CFG.min_lr)\n    elif CFG.scheduler == 'ReduceLROnPlateau':\n        scheduler = lr_scheduler.ReduceLROnPlateau(optimizer,\n                                                   mode='min',\n                                                   factor=0.1,\n                                                   patience=7,\n                                                   threshold=0.0001,\n                                                   min_lr=CFG.min_lr,)\n    elif CFG.scheduer == 'ExponentialLR':\n        scheduler = lr_scheduler.ExponentialLR(optimizer, gamma=0.85)\n    elif CFG.scheduler == None:\n        return None\n        \n    return scheduler","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:21.129480Z","iopub.execute_input":"2024-01-19T06:47:21.129824Z","iopub.status.idle":"2024-01-19T06:47:21.139175Z","shell.execute_reply.started":"2024-01-19T06:47:21.129795Z","shell.execute_reply":"2024-01-19T06:47:21.138113Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom collections import defaultdict\n\n# 定义优化器和学习率调度器\noptimizer = optim.Adam(model.parameters(), lr=CFG.lr, weight_decay=CFG.wd)\nscheduler = fetch_scheduler(optimizer)\n\nscheduler, optimizer","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:21.140401Z","iopub.execute_input":"2024-01-19T06:47:21.140757Z","iopub.status.idle":"2024-01-19T06:47:21.158373Z","shell.execute_reply.started":"2024-01-19T06:47:21.140730Z","shell.execute_reply":"2024-01-19T06:47:21.157171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 训练和验证函数","metadata":{}},{"cell_type":"code","source":"from tqdm import tqdm\nimport torch.cuda.amp as amp\nimport copy\n\n# 定义训练一个 epoch 的函数\ndef train_one_epoch(model, optimizer, scheduler, dataloader, device, epoch):\n    model.to(device)\n    # 设置模型为训练模式\n    model.train()\n    \n    # 使用混合精度训练\n    scaler = amp.GradScaler()\n    \n    dataset_size = 0\n    running_loss = 0.0\n    \n    # 使用 tqdm 显示训练进度\n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Train ')\n    \n    # 遍历数据加载器\n    for step, (images, masks) in pbar:         \n        images = images.to(device, dtype=torch.float)\n        masks  = masks.to(device, dtype=torch.long)\n        \n        batch_size = images.size(0)\n        \n        # 使用混合精度进行前向传播和计算损失\n        with amp.autocast(enabled=True):\n            y_pred = model(images)\n            loss   = criterion(y_pred, masks)\n            loss   = loss / CFG.n_accumulate\n            \n        # 使用混合精度进行反向传播\n        scaler.scale(loss).backward()\n    \n        if (step + 1) % CFG.n_accumulate == 0:\n            # 使用混合精度进行参数更新\n            scaler.step(optimizer)\n            scaler.update()\n\n            # zero the parameter gradients\n            optimizer.zero_grad()\n\n            # 如果使用学习率调度器，则进行调度\n            if scheduler is not None:\n                scheduler.step()\n                \n        # 计算累计损失和样本数\n        running_loss += (loss.item() * batch_size)\n        dataset_size += batch_size\n        \n        # 计算平均 epoch 损失\n        epoch_loss = running_loss / dataset_size\n        \n        # 在 tqdm 中更新训练进度显示\n        mem = torch.cuda.memory_reserved() / 1E9 if torch.cuda.is_available() else 0\n        current_lr = optimizer.param_groups[0]['lr']\n        pbar.set_postfix(train_loss=f'{epoch_loss:0.4f}',\n                        lr=f'{current_lr:0.5f}',\n                        gpu_mem=f'{mem:0.2f} GB')\n    \n    # 释放 GPU 内存\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    # 返回平均 epoch 损失\n    return epoch_loss","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:21.160035Z","iopub.execute_input":"2024-01-19T06:47:21.160373Z","iopub.status.idle":"2024-01-19T06:47:21.176574Z","shell.execute_reply.started":"2024-01-19T06:47:21.160344Z","shell.execute_reply":"2024-01-19T06:47:21.175368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test\n# train_one_epoch(model, optimizer, scheduler, train_loader, CFG.device, 1)","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:21.178075Z","iopub.execute_input":"2024-01-19T06:47:21.178424Z","iopub.status.idle":"2024-01-19T06:47:21.192036Z","shell.execute_reply.started":"2024-01-19T06:47:21.178395Z","shell.execute_reply":"2024-01-19T06:47:21.191068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef valid_one_epoch(model, dataloader, device, epoch):\n    model.to(device)\n    # 设置模型为评估模式\n    model.eval()\n    \n    dataset_size = 0\n    running_loss = 0.0\n    \n    val_scores = []\n    \n    # 使用 tqdm 显示验证进度\n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Valid ')\n    \n    # 遍历数据加载器\n    for step, (images, masks) in pbar:        \n        images  = images.to(device, dtype=torch.float)\n        masks   = masks.to(device, dtype=torch.long)\n        \n        batch_size = images.size(0)\n        \n        # 进行前向传播和计算损失\n        y_pred  = model(images)\n        loss    = criterion(y_pred, masks)\n        \n        # 计算累计损失和样本数\n        running_loss += (loss.item() * batch_size)\n        dataset_size += batch_size\n        \n        # 计算平均 epoch 损失\n        epoch_loss = running_loss / dataset_size\n        \n        # 计算accuracy\n        val_accuracy = calculate_accuracy(y_pred, masks)\n        val_scores.append(val_accuracy)\n        \n        # 在 tqdm 中更新验证进度显示\n        mem = torch.cuda.memory_reserved() / 1E9 if torch.cuda.is_available() else 0\n        current_lr = optimizer.param_groups[0]['lr']\n        pbar.set_postfix(valid_loss=f'{epoch_loss:0.4f}',\n                        lr=f'{current_lr:0.5f}',\n                        accuracy=f'{val_accuracy:0.3f}')\n    \n    # 计算accuracy\n    val_scores  = np.mean(val_scores)\n    \n    # 释放 GPU 内存\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    # 返回平均 epoch 损失和验证指标\n    return epoch_loss, val_scores","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:21.193418Z","iopub.execute_input":"2024-01-19T06:47:21.193771Z","iopub.status.idle":"2024-01-19T06:47:21.207909Z","shell.execute_reply.started":"2024-01-19T06:47:21.193742Z","shell.execute_reply":"2024-01-19T06:47:21.206820Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# valid_one_epoch(model, valid_loader, CFG.device, 1)","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:21.209311Z","iopub.execute_input":"2024-01-19T06:47:21.209667Z","iopub.status.idle":"2024-01-19T06:47:21.219324Z","shell.execute_reply.started":"2024-01-19T06:47:21.209626Z","shell.execute_reply":"2024-01-19T06:47:21.218169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# For colored terminal text\n# 从 colorama 库中导入 Fore（前景色）、Back（背景色）和 Style（样式）模块\nfrom colorama import Fore, Back, Style\nc_  = Fore.GREEN\nsr_ = Style.RESET_ALL\n\ndef run_training(model, optimizer, scheduler, device, num_epochs):\n    # 通过 wandb.watch 自动记录梯度\n    wandb.watch(model, log_freq=100)\n    \n    # 如果有可用的 CUDA 设备，打印设备信息\n    if torch.cuda.is_available():\n        print(\"cuda: {}\\n\".format(torch.cuda.get_device_name()))\n    \n    # 记录训练过程的变量\n    start = time.time()\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_accuracy      = -np.inf\n    best_epoch     = -1\n    history = defaultdict(list)\n    \n    # 遍历每个 epoch\n    for epoch in range(1, num_epochs + 1): \n        gc.collect()\n        print(f'Epoch {epoch}/{num_epochs}', end='')\n        \n        # 训练一个 epoch，并记录训练损失\n        train_loss = train_one_epoch(model, optimizer, scheduler, \n                                     dataloader=train_loader, \n                                     device=CFG.device, epoch=epoch)\n        \n        # 验证一个 epoch，并记录验证损失和指标\n        val_loss, val_scores = valid_one_epoch(model, valid_loader, \n                                               device=CFG.device, \n                                               epoch=epoch)\n    \n        # 记录训练和验证指标\n        history['Train Loss'].append(train_loss)\n        history['Valid Loss'].append(val_loss)\n        history['Valid accuracy'].append(val_scores)\n        \n        # 使用 wandb 记录训练和验证指标\n        wandb.log({\"Train Loss\": train_loss, \n                   \"Valid Loss\": val_loss,\n                   \"Valid Accuracy\": val_scores,\n                   \"LR\":scheduler.get_last_lr()[0]})\n        \n        # 打印验证 Dice 和 Jaccard 指标\n        print(f'Valid accuracy: {val_scores:0.4f}')\n        \n        # 更新最佳模型权重\n        if val_scores >= best_accuracy:\n            print(f\"{c_}Valid Score Improved ({best_accuracy:0.4f} ---> {val_scores:0.4f})\")\n            best_accuracy    = val_scores\n            \n            best_epoch   = epoch\n            run.summary[\"Best Accuracy\"]    = best_accuracy\n            run.summary[\"Best Epoch\"]   = best_epoch\n            best_model_wts = copy.deepcopy(model.state_dict())\n            PATH = f\"best_epoch-{fold:02d}.bin\"\n            torch.save(model.state_dict(), PATH)\n            # 保存模型文件到当前目录\n            wandb.save(PATH)\n            print(f\"Model Saved{sr_}\")\n            \n        # 保存当前模型权重\n        last_model_wts = copy.deepcopy(model.state_dict())\n        PATH = f\"last_epoch-{fold:02d}.bin\"\n        torch.save(model.state_dict(), PATH)\n            \n        print()\n    \n    # 训练完成，计算总体时间\n    end = time.time()\n    time_elapsed = end - start\n    print('Training complete in {:.0f}h {:.0f}m {:.0f}s'.format(\n        time_elapsed // 3600, (time_elapsed % 3600) // 60, (time_elapsed % 3600) % 60))\n    print(\"Best accuracy: {:.4f}\".format(best_accuracy))\n    \n    # 加载最佳模型权重\n    model.load_state_dict(best_model_wts)\n    \n    # 返回训练完成的模型和训练历史记录\n    return model, history","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:21.221114Z","iopub.execute_input":"2024-01-19T06:47:21.221424Z","iopub.status.idle":"2024-01-19T06:47:21.238315Z","shell.execute_reply.started":"2024-01-19T06:47:21.221397Z","shell.execute_reply":"2024-01-19T06:47:21.237106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 训练","metadata":{}},{"cell_type":"code","source":"import wandb\n# 尝试从 Kaggle secrets 中获取 Weights & Biases 的 API 密钥进行登录\ntry:\n    # 从 kaggle_secrets 模块导入 UserSecretsClient 类\n    from kaggle_secrets import UserSecretsClient\n    # 创建 UserSecretsClient 实例\n    user_secrets = UserSecretsClient()\n    # 获取名为 \"WANDB\" 的密钥（Weights & Biases API 密钥）\n    api_key = user_secrets.get_secret(\"WANDB\")\n    # 使用获取的 API 密钥登录到 Weights & Biases\n    wandb.login(key=api_key)\n    # 如果登录成功，将 anonymous 变量设为 None\n    anonymous = None\nexcept:\n    # 如果发生异常，将 anonymous 变量设为 \"must\"，并输出提示信息\n    anonymous = \"must\"\n    print('要使用您的 W&B 账户，\\n请转到 Add-ons -> Secrets 并提供您的 W&B 访问令牌。使用标签名为 WANDB。\\n从此处获取您的 W&B 访问令牌: https://wandb.ai/authorize')","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:21.245292Z","iopub.execute_input":"2024-01-19T06:47:21.245676Z","iopub.status.idle":"2024-01-19T06:47:41.291217Z","shell.execute_reply.started":"2024-01-19T06:47:21.245643Z","shell.execute_reply":"2024-01-19T06:47:41.290060Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def main():\n    for fold in range(1):\n        print(f'#'*15)\n        print(f'### Fold: {fold}')\n        print(f'#'*15)\n\n        # 初始化 wandb，并设置项目、配置、匿名等信息\n        run = wandb.init(project='leaf_disease_class', \n                         config={k:v for k, v in dict(vars(CFG)).items() if '__' not in k},\n                         anonymous=anonymous,\n                         name=f\"fold-{fold}|dim-{CFG.img_size[0]}x{CFG.img_size[1]}|model-{CFG.model_name}\",\n                         group=CFG.comment,\n                        )\n\n        # 准备训练和验证数据加载器\n        train_loader, valid_loader = prepare_loaders(fold=fold)\n\n        # 构建模型\n        model = model.to(CFG.device)\n\n\n        # 进行模型训练\n        model, history = run_training(model, optimizer, scheduler,\n                                      device=CFG.device,\n                                      num_epochs=CFG.epochs)\n\n        # 结束 wandb 实验\n        run.finish()","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:41.293897Z","iopub.execute_input":"2024-01-19T06:47:41.294347Z","iopub.status.idle":"2024-01-19T06:47:41.303716Z","shell.execute_reply.started":"2024-01-19T06:47:41.294306Z","shell.execute_reply":"2024-01-19T06:47:41.302658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 预测","metadata":{}},{"cell_type":"code","source":"def load_model(model, weight_path):\n    model.load_state_dict(torch.load(weight_path))\n    \n    model.eval()\n    model.to(CFG.device)\n    return model","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:41.305160Z","iopub.execute_input":"2024-01-19T06:47:41.305867Z","iopub.status.idle":"2024-01-19T06:47:41.317485Z","shell.execute_reply.started":"2024-01-19T06:47:41.305812Z","shell.execute_reply":"2024-01-19T06:47:41.316265Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_submisssion = pd.read_csv('/kaggle/input/cassava-leaf-disease-classification/sample_submission.csv')\nsample_submisssion.head()","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:41.318824Z","iopub.execute_input":"2024-01-19T06:47:41.319252Z","iopub.status.idle":"2024-01-19T06:47:41.345158Z","shell.execute_reply.started":"2024-01-19T06:47:41.319214Z","shell.execute_reply":"2024-01-19T06:47:41.344038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ds = BuildDataset(sample_submisssion, False, transforms=data_transforms['valid'])\n\ntest_loader = DataLoader(test_ds, batch_size=CFG.valid_bs, \n                              num_workers=4, shuffle=False, pin_memory=True)\n\nimgs = next(iter(test_loader))\nimgs.size()","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:41.346536Z","iopub.execute_input":"2024-01-19T06:47:41.346953Z","iopub.status.idle":"2024-01-19T06:47:41.522710Z","shell.execute_reply.started":"2024-01-19T06:47:41.346915Z","shell.execute_reply":"2024-01-19T06:47:41.521313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"weight_path = '/kaggle/input/leaf-disease-classifiction-weight/best_epoch-00.bin'\nmodel = load_model(model, weight_path)\npreds_list = []\n\nfor imgs in test_loader:\n    imgs = imgs.to(CFG.device)\n    \n    preds = model(imgs)\n    preds = torch.argmax(preds, dim=1)\n    preds = preds.detach().cpu().numpy()\n    preds_list.extend(preds)\n\npreds_list","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:41.524732Z","iopub.execute_input":"2024-01-19T06:47:41.525648Z","iopub.status.idle":"2024-01-19T06:47:41.926612Z","shell.execute_reply.started":"2024-01-19T06:47:41.525597Z","shell.execute_reply":"2024-01-19T06:47:41.925228Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imgs = imgs.cpu().detach()\nplot_batch(imgs, preds_list, 1)","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:41.928687Z","iopub.execute_input":"2024-01-19T06:47:41.929555Z","iopub.status.idle":"2024-01-19T06:47:42.198538Z","shell.execute_reply.started":"2024-01-19T06:47:41.929507Z","shell.execute_reply":"2024-01-19T06:47:42.197450Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imgs_list = os.listdir(CFG.test_img_dir)\nimgs_list","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:42.200176Z","iopub.execute_input":"2024-01-19T06:47:42.201066Z","iopub.status.idle":"2024-01-19T06:47:42.208245Z","shell.execute_reply.started":"2024-01-19T06:47:42.201027Z","shell.execute_reply":"2024-01-19T06:47:42.207026Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sub = pd.DataFrame({'image_id':imgs_list, 'label':preds_list})\ndf_sub.head()","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:42.209767Z","iopub.execute_input":"2024-01-19T06:47:42.210206Z","iopub.status.idle":"2024-01-19T06:47:42.225373Z","shell.execute_reply.started":"2024-01-19T06:47:42.210169Z","shell.execute_reply":"2024-01-19T06:47:42.224259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sub.to_csv('submission.csv', index=False)\nprint('submission success')","metadata":{"execution":{"iopub.status.busy":"2024-01-19T06:47:42.226817Z","iopub.execute_input":"2024-01-19T06:47:42.227247Z","iopub.status.idle":"2024-01-19T06:47:42.237444Z","shell.execute_reply.started":"2024-01-19T06:47:42.227210Z","shell.execute_reply":"2024-01-19T06:47:42.236335Z"},"trusted":true},"execution_count":null,"outputs":[]}]}