{"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"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":5048,"databundleVersionId":868335,"sourceType":"competition"}],"dockerImageVersionId":30558,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# [CNN For State Farm Distracted Driver Detection](https://www.kaggle.com/superbigfive/cnn-for-state-farm-distracted-driver-detection)\n- 作者 : 常铭，电子科技大学，计算机科学与工程学院 (网络空间安全学院)；\n- 时间 : 2023-10-26;\n- 主要内容：基于 CNN 的司机疲劳驾驶行为检测；\n- 本实验在 Kaggle 网站上运行，深度学习框架为 Pytorch，从标题链接中 fork 一份即可一键运行；\n- 数据集选取 Kaggle 某场比赛中提供的[数据集](https://www.kaggle.com/competitions/state-farm-distracted-driver-detection)，共有十种驾驶行为，其中有九种为疲劳驾驶行为，一种为正常驾驶行为。训练集中每种驾驶行为对应的图片有 2k+ 张，测试集大概 80k 张图片；\n- 网络模型使用的 CNN，网络结构由三个卷积块和三个全连接块构成，添加了 Dropout 层，减缓模型过拟合现象；\n- 损失函数设置为交叉熵损失函数，优化器选取 Adam，学习率和 weight_decay 分别设置为 0.01 和 1e-6 ，利用了 CosineAnnealingLR 学习率调度器实现训练过程中学习率的动态调整。","metadata":{"papermill":{"duration":0.011231,"end_time":"2023-10-27T00:44:20.934818","exception":false,"start_time":"2023-10-27T00:44:20.923587","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# 📚 Import Libraries\n导入本次实验所需要的库。 <br/>","metadata":{"papermill":{"duration":0.010332,"end_time":"2023-10-27T00:44:20.955937","exception":false,"start_time":"2023-10-27T00:44:20.945605","status":"completed"},"tags":[]}},{"cell_type":"code","source":"%load_ext autoreload\n%autoreload 2","metadata":{"papermill":{"duration":0.047684,"end_time":"2023-10-27T00:44:21.014166","exception":false,"start_time":"2023-10-27T00:44:20.966482","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-10-27T08:56:35.097073Z","iopub.execute_input":"2023-10-27T08:56:35.097701Z","iopub.status.idle":"2023-10-27T08:56:35.127219Z","shell.execute_reply.started":"2023-10-27T08:56:35.097661Z","shell.execute_reply":"2023-10-27T08:56:35.126508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torch.optim import lr_scheduler\nimport torchvision\nfrom torch.cuda import amp\n\n# 绘图相关\nfrom matplotlib import pyplot as plt\nfrom pylab import mpl\n# 指定默认字体：解决plot不能显示中文问题\nmpl.rcParams['font.sans-serif'] = ['Microsoft YaHei'] \n# 解决保存图像是负号'-'显示为方块的问题\nmpl.rcParams['axes.unicode_minus'] = False \n\nimport os\nimport numpy as np\nimport copy, random, time\n\nimport gc\n\n# 数据及相关\nimport cv2\nimport pandas as pd\nfrom PIL import Image\nimport albumentations as A\nfrom sklearn.model_selection import train_test_split\n\nfrom IPython import display as ipd\nfrom tqdm.notebook import tqdm\n# 使得终端可以打印不同颜色的文字\nfrom colorama import Fore, Back, Style\nc_  = Fore.GREEN\nsr_ = Style.RESET_ALL\nfrom collections import defaultdict\n\nimport warnings\nwarnings.filterwarnings (\"ignore\")","metadata":{"papermill":{"duration":7.127938,"end_time":"2023-10-27T00:44:28.152971","exception":false,"start_time":"2023-10-27T00:44:21.025033","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-10-27T08:56:35.129038Z","iopub.execute_input":"2023-10-27T08:56:35.129579Z","iopub.status.idle":"2023-10-27T08:56:40.322436Z","shell.execute_reply.started":"2023-10-27T08:56:35.129547Z","shell.execute_reply":"2023-10-27T08:56:40.321454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ⭐ WandB\nWeights & Bias (W&B) 是一个可以跟踪实验的平台，利用该平台可以很方便地可视化实验过程中各项指标，进行对比分析，便于我们调整模型的各项参数。","metadata":{"papermill":{"duration":0.010431,"end_time":"2023-10-27T00:44:28.174066","exception":false,"start_time":"2023-10-27T00:44:28.163635","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!pip install -qq -U wandb\nimport wandb\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient ()\napi_key = user_secrets.get_secret (\"WANDB\")\nwandb.login (key = api_key)\nanonymous = None","metadata":{"papermill":{"duration":22.226793,"end_time":"2023-10-27T00:44:50.411493","exception":false,"start_time":"2023-10-27T00:44:28.184700","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-10-27T08:56:40.323634Z","iopub.execute_input":"2023-10-27T08:56:40.324042Z","iopub.status.idle":"2023-10-27T08:56:59.360558Z","shell.execute_reply.started":"2023-10-27T08:56:40.324016Z","shell.execute_reply":"2023-10-27T08:56:59.359819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ⚙️ Configuration\n本实验中所有用到的参数都将统一设置在 CFG 类中，便于查看与修改。","metadata":{"papermill":{"duration":0.010689,"end_time":"2023-10-27T00:44:50.433715","exception":false,"start_time":"2023-10-27T00:44:50.423026","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class CFG :\n    seed          = 20231026\n    exp_name      = 'CNN For State Farm Distracted Driver Detection'\n    model_name    = 'cnn'\n    train_bs      = 32\n    valid_bs      = 64\n    test_bs       = 64\n    num_workers   = 2\n    img_size      = [96, 128]\n    epochs        = 25\n    lr            = 0.01\n    wd            = 1e-6\n    scheduler     = 'CosineAnnealingLR'\n    optimizer     = 'Adam'\n    min_lr        = 1e-6\n    T_max         = int (30000 / train_bs * epochs) + 50\n    T_0           = 25\n    warmup_epochs = 0\n    n_accumulate  = max (1, 32 // train_bs)\n    num_classes   = 10\n    device        = torch.device (\"cuda:0\" if torch.cuda.is_available() else \"cpu\")","metadata":{"papermill":{"duration":0.086531,"end_time":"2023-10-27T00:44:50.531303","exception":false,"start_time":"2023-10-27T00:44:50.444772","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-10-27T08:56:59.362532Z","iopub.execute_input":"2023-10-27T08:56:59.362960Z","iopub.status.idle":"2023-10-27T08:56:59.448605Z","shell.execute_reply.started":"2023-10-27T08:56:59.362934Z","shell.execute_reply":"2023-10-27T08:56:59.447628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ❗ Reproducibility\n这里设置了所有随机种子的值，保证了每次运行时，运行结果都是相同的。","metadata":{"papermill":{"duration":0.010612,"end_time":"2023-10-27T00:44:50.553157","exception":false,"start_time":"2023-10-27T00:44:50.542545","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def set_seed (seed = CFG.seed):\n    '''\n        初始化各个随机种子为同一值, \n        保证每次执行程序的运行情况都相同\n    '''\n    np.random.seed (seed)\n    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 = False\n    os.environ['PYTHONHASHSEED'] = str (seed)\n    print ('> SEEDING DONE')\n    \nset_seed (CFG.seed)","metadata":{"papermill":{"duration":0.08997,"end_time":"2023-10-27T00:44:50.653981","exception":false,"start_time":"2023-10-27T00:44:50.564011","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-10-27T08:56:59.449723Z","iopub.execute_input":"2023-10-27T08:56:59.450073Z","iopub.status.idle":"2023-10-27T08:56:59.516138Z","shell.execute_reply.started":"2023-10-27T08:56:59.450042Z","shell.execute_reply":"2023-10-27T08:56:59.515214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🔨 Utility\n基于本实验设置的一些简单的工具。","metadata":{"papermill":{"duration":0.010572,"end_time":"2023-10-27T00:44:50.675622","exception":false,"start_time":"2023-10-27T00:44:50.665050","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def showImage (imgs, labels, size = 5) :\n    '''\n        输入是若干张图片组成的一个 list\n        每张图片的格式都是张量，数据范围是 [0, 1]\n    '''\n    size = min (size, len (imgs))\n    plt.figure (figsize = (5 * size, 5))\n    for i in range (size) :\n        plt.subplot (1, size, i + 1)\n        img = imgs[i, ].permute ((1, 2, 0)).detach ().cpu ().numpy () * 255.0\n        img = img.astype ('uint8')\n        plt.axis ('off')\n        plt.title (num2label[int (labels[i])], fontsize = 'xx-large')\n        plt.imshow (img)\n    plt.tight_layout ()\n    plt.show ()\n    ","metadata":{"papermill":{"duration":0.083954,"end_time":"2023-10-27T00:44:50.770634","exception":false,"start_time":"2023-10-27T00:44:50.686680","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-10-27T08:56:59.517127Z","iopub.execute_input":"2023-10-27T08:56:59.517429Z","iopub.status.idle":"2023-10-27T08:56:59.579517Z","shell.execute_reply.started":"2023-10-27T08:56:59.517405Z","shell.execute_reply":"2023-10-27T08:56:59.578687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🍚 Dataset\n创建一个继承 $Dataset$ 类的子类, 重写 $\\_ \\_ len \\_ \\_$ 和 $\\_ \\_ getitem \\_ \\_$ 方法, 前者用来获取数据及大小, 后者用于指定下标, 返回对应的样本。<br/>\n$BuildDataset$ 类中存储了每一个图像的文件路径 ($self.img_paths$) 及其标签 ($self.labels$)，从数据集中取出数据时，根据数据对应的文件路径，利用 $Image.open$ 方法读取图像，并转化为 ”$RGB$“ 图像。转化成 $numpy$ 类型的数据后再经 $transform$ 实现对图像的数据增强。<br/>\n数据集可被划分为训练集、验证集、测试集，对于训练集和验证集，它们的图像数据对应的文件路径相同，测试集对应的图想保存在另一个文件夹中。","metadata":{"papermill":{"duration":0.01097,"end_time":"2023-10-27T00:44:50.792889","exception":false,"start_time":"2023-10-27T00:44:50.781919","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class BuildDataset (torch.utils.data.Dataset) :\n    def __init__ (self, Datatype, transforms = None) :\n        self.transforms = transforms\n        self.Datatype   = Datatype\n        \n        self.img_paths = []\n        self.labels    = []\n        \n        if self.Datatype is not \"test\" :\n            BASE_PATH  = \"../input/state-farm-distracted-driver-detection/imgs/train\"\n            for i in range (10) :\n                img_dir = f\"{BASE_PATH}/c{i}\"\n                for img_name in os.listdir (img_dir) :\n                    self.img_paths.append (f\"{img_dir}/{img_name}\")\n                    self.labels.append (i)\n        else :\n            BASE_PATH = \"../input/state-farm-distracted-driver-detection/imgs/test\"\n            for img_name in os.listdir (BASE_PATH) :\n                self.img_paths.append (f\"{BASE_PATH}/{img_name}\")\n        \n    def __len__ (self) :\n        return len (self.img_paths)\n    \n    def __getitem__ (self, idx) :\n        img_path = self.img_paths[idx]\n        img = Image.open (img_path).convert (\"RGB\")\n        img = np.asarray (img) / 255.0\n        if self.transforms :\n            data = self.transforms[self.Datatype] (image = img)\n            img  = data['image']\n        img = np.transpose (img, (2, 0, 1))\n        if self.Datatype is not \"test\" :\n            label = self.labels[idx]\n            return torch.tensor (img), torch.tensor (label)\n        else :\n            return torch.tensor (img)","metadata":{"papermill":{"duration":0.086899,"end_time":"2023-10-27T00:44:50.890852","exception":false,"start_time":"2023-10-27T00:44:50.803953","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-10-27T08:56:59.580593Z","iopub.execute_input":"2023-10-27T08:56:59.580880Z","iopub.status.idle":"2023-10-27T08:56:59.641439Z","shell.execute_reply.started":"2023-10-27T08:56:59.580853Z","shell.execute_reply":"2023-10-27T08:56:59.640608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🌈 Augmentations\n本实验中并未对训练集进行图像增强。经过本人多次实验，发现在执行图像增强操作后模型的精度会下降，而不执行图像增强操作时，模型并不存在过拟合的现象。验证集的准确率高达 $99\\%$ 以上。这很令人感到困惑，因为训练过程中模型在训练集上的准确率仅有 $82\\%$，而验证集上的效果却高了很多；在测试集上测试模型的分类效果，发现偶尔会出现分类错误的情况，但总的准确率一定比 $99\\%$ 要低。目前不太清楚原因，在 $Kaggle$ 上阅读他人分享的 $notebook$ 时发现也存在“模型在验证集上分类效果要比训练集好“的情况，但是相对来说本人实验中该现象更明显一些。","metadata":{"papermill":{"duration":0.010727,"end_time":"2023-10-27T00:44:50.912952","exception":false,"start_time":"2023-10-27T00:44:50.902225","status":"completed"},"tags":[]}},{"cell_type":"code","source":"data_transforms = {\n    \"train\": A.Compose ([\n        # 固定图像尺寸\n        A.Resize (*CFG.img_size, interpolation = cv2.INTER_NEAREST),\n        ], p = 1.0),\n    \n    \"valid\": A.Compose ([\n        A.Resize (*CFG.img_size, interpolation = cv2.INTER_NEAREST),\n        ], p = 1.0),\n    \n    \"test\": A.Compose ([\n        A.Resize (*CFG.img_size, interpolation = cv2.INTER_NEAREST),\n        ], p = 1.0)\n}","metadata":{"papermill":{"duration":0.082718,"end_time":"2023-10-27T00:44:51.006679","exception":false,"start_time":"2023-10-27T00:44:50.923961","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-10-27T08:56:59.642482Z","iopub.execute_input":"2023-10-27T08:56:59.642777Z","iopub.status.idle":"2023-10-27T08:56:59.698755Z","shell.execute_reply.started":"2023-10-27T08:56:59.642752Z","shell.execute_reply":"2023-10-27T08:56:59.697929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🍰 DataLoader\n利用 train_test_split 方法按照 9 : 1 的比例将带标签的数据集划分为训练集和验证集，没有标签的数据集全部作为测试集；<br/>\n然后利用 torch.utils.data.DataLoader 方法将训练集、验证集、测试集打包成若干个固定大小的批次。","metadata":{"papermill":{"duration":0.010665,"end_time":"2023-10-27T00:44:51.028497","exception":false,"start_time":"2023-10-27T00:44:51.017832","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def DataLoader (data_transforms) :\n    train_and_test_dataset = BuildDataset (\"train\", transforms = data_transforms)\n    indexes = list (range (len (train_and_test_dataset)))\n    train_indexes, valid_indexes = train_test_split (indexes, test_size = 0.1)\n    train_dataset = torch.utils.data.Subset (train_and_test_dataset, train_indexes)\n    valid_dataset = torch.utils.data.Subset (train_and_test_dataset, valid_indexes)\n    test_dataset  = BuildDataset (\"test\", transforms = data_transforms)\n    valid_dataset.Datatype = \"valid\"\n    \n    # pin_memory 参数表示是否将加载的数据常驻内存\n    # drop_last 参数表示是否丢弃最后一个批次 (可能不满 batch_size 个样本)\n    train_loader = torch.utils.data.DataLoader (train_dataset, batch_size = CFG.train_bs, num_workers = CFG.num_workers, \n                                                shuffle = True, pin_memory = True, drop_last = False)\n    valid_loader = torch.utils.data.DataLoader (valid_dataset, batch_size = CFG.valid_bs, num_workers = CFG.num_workers, \n                                                shuffle = False, pin_memory = True, drop_last = False)\n    test_loader  = torch.utils.data.DataLoader (test_dataset, batch_size = CFG.test_bs, num_workers = CFG.num_workers,\n                                                shuffle = False, pin_memory = True, drop_last = False)\n    \n    return train_loader, valid_loader, test_loader\n\nnum2label = {0 : 'Safe driving',                 1 : 'Texting - right', \n             2 : 'Talking on the phone - right', 3 : 'Texting - left', \n             4 : 'Talking on the phone - left',  5 : 'Operating the radio', \n             6 : 'Drinking',                     7 : 'Reaching behind', \n             8 : 'Hair and makeup',              9 : 'Talking to passenger'}\n\ntrain_loader, valid_loader, test_loader = DataLoader (data_transforms)\n\nprint (\"len of batch : \", len (train_loader))\nit = iter (train_loader)\nimgs, labels = next (it)\nprint (imgs.size (), labels.size ())","metadata":{"papermill":{"duration":3.362715,"end_time":"2023-10-27T00:44:54.402299","exception":false,"start_time":"2023-10-27T00:44:51.039584","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-10-27T08:56:59.699746Z","iopub.execute_input":"2023-10-27T08:56:59.700031Z","iopub.status.idle":"2023-10-27T08:57:09.904722Z","shell.execute_reply.started":"2023-10-27T08:56:59.700008Z","shell.execute_reply":"2023-10-27T08:57:09.902821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 📈 Visualization\n随机展示部分图像，检查数据集的读入是否出错、检测数据增强效果。<br/>\n为了避免计算量过大，强制将图像缩小为 $96 \\times 128$ 的大小，可以明显看出图像的低分辨率。<br/>\n实验证明，即使这么做会损失部分信息，但模型最终依然能够高效地分类疲劳驾驶行为。","metadata":{"papermill":{"duration":0.01412,"end_time":"2023-10-27T00:44:54.430838","exception":false,"start_time":"2023-10-27T00:44:54.416718","status":"completed"},"tags":[]}},{"cell_type":"code","source":"showImage (imgs, labels)","metadata":{"papermill":{"duration":0.89534,"end_time":"2023-10-27T00:44:55.340195","exception":false,"start_time":"2023-10-27T00:44:54.444855","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-10-27T08:57:09.911261Z","iopub.execute_input":"2023-10-27T08:57:09.911732Z","iopub.status.idle":"2023-10-27T08:57:10.646461Z","shell.execute_reply.started":"2023-10-27T08:57:09.911684Z","shell.execute_reply":"2023-10-27T08:57:10.645586Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 📦 Model\n由于网络结构比较简单，所以手写实现了该网络，先实现了卷积块和全连接块，然后再实现了 $CNN$。","metadata":{"papermill":{"duration":0.01529,"end_time":"2023-10-27T00:44:55.371274","exception":false,"start_time":"2023-10-27T00:44:55.355984","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class ConvBlock (nn.Module) :\n    def __init__ (self, in_channels, out_channels) :\n        super (ConvBlock, self).__init__ ()\n        self.main = nn.Sequential (\n            nn.Conv2d (in_channels, out_channels, 3, 1, 1, bias = False),\n            nn.ReLU (True),\n            nn.BatchNorm2d (out_channels),\n            nn.Conv2d (out_channels, out_channels, 3, 1, 1, bias = False),\n            nn.ReLU (True),\n            nn.BatchNorm2d (out_channels),\n            nn.MaxPool2d (2),\n            nn.Dropout (p = 0.3)\n        )\n\n    def forward (self, x) :\n        x = self.main (x)\n        return x\n        \nclass DenseBlock (nn.Module) :\n    def __init__ (self, in_channels, out_channels,) :\n        super (DenseBlock, self).__init__ ()\n        self.main = nn.Sequential (\n            nn.Linear (in_channels, out_channels),\n            nn.BatchNorm1d (out_channels),\n            nn.Dropout (p = 0.25)\n        )\n        \n    def forward (self, x) :\n        x = self.main (x)\n        return x\n        \nclass CNN (nn.Module) :\n    '''\n        W' = (W - F + 2 * P) / S + 1\n        input_size = (b, 3, 96, 128)\n        conv1 : (b, 32, 48, 64)\n        conv2 : (b, 64, 24, 32)\n        conv3 : (b, 128, 12, 16), 24576\n        dense1 : (b, 512)\n        dense2 : (b, 128)\n        dense3 : (b, 10)\n        output_size = (b, 10)\n    '''\n    def __init__ (self, ConvBlock, DenseBlock) :\n        super (CNN, self).__init__ ()\n        self.main = nn.Sequential (\n            ConvBlock (3, 32),\n            ConvBlock (32, 64),\n            ConvBlock (64, 128),\n            nn.Flatten (),\n            DenseBlock (128 * 12 * 16, 512),\n            DenseBlock (512, 128),\n            DenseBlock (128, 10),\n            nn.Softmax ()\n        )\n    \n    def forward (self, x) :\n        x = self.main (x)\n        return x\n    \n\ndef LoadModel (LOADPATH) :\n    '''\n        用于测试模型前加载训练过程中表现最好的模型\n    '''\n    model = CNN ()\n    model.load_state_dict (torch.load (LOADPATH))\n    model.eval ()\n    return model","metadata":{"papermill":{"duration":0.094441,"end_time":"2023-10-27T00:44:55.481287","exception":false,"start_time":"2023-10-27T00:44:55.386846","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-10-27T08:57:10.647592Z","iopub.execute_input":"2023-10-27T08:57:10.647853Z","iopub.status.idle":"2023-10-27T08:57:10.712473Z","shell.execute_reply.started":"2023-10-27T08:57:10.647829Z","shell.execute_reply":"2023-10-27T08:57:10.711692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🔧 Loss Function\n损失函数使用交叉熵损失函数，同时实现了计算准确个数的函数。","metadata":{"papermill":{"duration":0.015297,"end_time":"2023-10-27T00:44:55.513082","exception":false,"start_time":"2023-10-27T00:44:55.497785","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def calc_correct (preds, labels) :\n    _, preds = torch.max (preds, 1)\n    return (preds == labels).float ().sum ().cpu ().item ()\n\ndef criterion (preds, labels) :\n    return nn.CrossEntropyLoss () (preds, labels).sum ()","metadata":{"papermill":{"duration":0.086507,"end_time":"2023-10-27T00:44:55.615672","exception":false,"start_time":"2023-10-27T00:44:55.529165","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-10-27T08:57:10.713634Z","iopub.execute_input":"2023-10-27T08:57:10.713973Z","iopub.status.idle":"2023-10-27T08:57:10.770962Z","shell.execute_reply.started":"2023-10-27T08:57:10.713939Z","shell.execute_reply":"2023-10-27T08:57:10.770007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🔍 Optimizer\ntorch.optim.CosineAnnealingLR 方法是一个用于设置学习率的调度器，它可以根据一个余弦退火的策略，动态地调整优化器的学习率。它的原理是在每个周期内，将学习率从一个最大值降低到一个最小值，然后在下一个周期内重复这个过程。这样可以避免学习率过大或过小导致的收敛困难或局部最优;","metadata":{"papermill":{"duration":0.01517,"end_time":"2023-10-27T00:44:55.646365","exception":false,"start_time":"2023-10-27T00:44:55.631195","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def fetch_scheduler (optimizer):\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":{"papermill":{"duration":0.087006,"end_time":"2023-10-27T00:44:55.748819","exception":false,"start_time":"2023-10-27T00:44:55.661813","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-10-27T08:57:10.772062Z","iopub.execute_input":"2023-10-27T08:57:10.772642Z","iopub.status.idle":"2023-10-27T08:57:10.828793Z","shell.execute_reply.started":"2023-10-27T08:57:10.772606Z","shell.execute_reply":"2023-10-27T08:57:10.827901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🚄 Training Function\n每一个 epoch 训练过程的具体实现。模型训练采用了自动混合精度策略。<br/>","metadata":{"papermill":{"duration":0.01652,"end_time":"2023-10-27T00:44:55.781623","exception":false,"start_time":"2023-10-27T00:44:55.765103","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def train_one_epoch (model, optimizer, scheduler, dataloader, device, epoch) :\n    model.train ()\n    # 创建一个 GradScaler 对象, 可以在迭代过程中动态估计损失放大的倍数\n    scaler = amp.GradScaler ()\n    \n    dataset_size = 0\n    running_loss, num_correct = 0.0, 0\n    epoch_loss, epoch_accuracy = 0.0, 0.0\n    \n    # 将一个可迭代对象作为参数传入，然后返回一个包装后的可迭代对象，\n    # 可以像平常一样对其进行迭代，每次请求一个值时，都会打印一个进度条。\n    pbar = tqdm (enumerate (dataloader), total = len (dataloader), desc = 'Train ')\n    for step, (images, labels) in pbar:         \n        images = images.to (device, dtype = torch.float)\n        labels = labels.to (device)\n        batch_size = images.size (0)\n        \n        # 前向传播过程中自动混合精度训练\n        with amp.autocast (enabled = True):\n            y_pred    = model (images)\n            y_pred    = y_pred.to (device, dtype = torch.float)\n            loss      = criterion (y_pred, labels)\n            \n        running_loss += loss.item ()\n        num_correct  += calc_correct (y_pred, labels)\n        # 放大损失、反向传播\n        scaler.scale (loss).backward ()\n        if (step + 1) % CFG.n_accumulate == 0 :\n            # 根据原放大倍数，梯度更新时缩小相应的倍数\n            scaler.step (optimizer)\n            # 更新损失放大的倍数\n            scaler.update ()\n            optimizer.zero_grad ()\n            if scheduler is not None :\n                # 更新学习率\n                scheduler.step ()\n        \n        dataset_size += batch_size\n        \n        epoch_loss     = running_loss / dataset_size\n        epoch_accuracy = num_correct / dataset_size\n        \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                          train_accuracy = f'{epoch_accuracy : 0.4f}',\n                          lr = f'{current_lr : 0.5f}',\n                          gpu_mem = f'{mem : 0.2f} GB')\n    torch.cuda.empty_cache ()\n    gc.collect ()\n    \n    return epoch_loss","metadata":{"papermill":{"duration":0.091475,"end_time":"2023-10-27T00:44:55.889136","exception":false,"start_time":"2023-10-27T00:44:55.797661","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-10-27T08:57:10.830131Z","iopub.execute_input":"2023-10-27T08:57:10.830400Z","iopub.status.idle":"2023-10-27T08:57:10.891394Z","shell.execute_reply.started":"2023-10-27T08:57:10.830377Z","shell.execute_reply":"2023-10-27T08:57:10.890610Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 👀 Validation Function\n评估函数的实现，使用验证集评估模型的分类效果和泛化性能。","metadata":{"papermill":{"duration":0.015078,"end_time":"2023-10-27T00:44:55.919807","exception":false,"start_time":"2023-10-27T00:44:55.904729","status":"completed"},"tags":[]}},{"cell_type":"code","source":"@torch.no_grad ()\ndef valid_one_epoch (model, dataloader, device, epoch):\n    model.eval ()\n    \n    dataset_size = 0\n    running_loss, num_correct = 0.0, 0\n    epoch_loss, epoch_accuracy = 0.0, 0.0\n    \n    pbar = tqdm (enumerate (dataloader), total = len (dataloader), desc = 'Valid ')\n    for step, (images, labels) in pbar :\n        images  = images.to (device, dtype = torch.float)\n        labels  = labels.to (device)\n        batch_size = images.size (0)\n        \n        y_pred    = model (images)\n        y_pred    = y_pred.to (device, dtype = torch.float)\n        loss      = criterion (y_pred, labels)\n        \n        running_loss += loss.item ()\n        num_correct  += calc_correct (y_pred, labels)\n        dataset_size += batch_size\n        \n        epoch_loss     = running_loss / dataset_size\n        epoch_accuracy = num_correct / dataset_size\n        \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                          valid_accuracy = f'{epoch_accuracy : 0.4f}',\n                          lr = f'{current_lr : 0.5f}',\n                          gpu_mem = f'{mem : 0.2f} GB')\n    torch.cuda.empty_cache ()\n    gc.collect ()\n    \n    return epoch_loss, epoch_accuracy","metadata":{"papermill":{"duration":0.088403,"end_time":"2023-10-27T00:44:56.024329","exception":false,"start_time":"2023-10-27T00:44:55.935926","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-10-27T08:57:10.892502Z","iopub.execute_input":"2023-10-27T08:57:10.892834Z","iopub.status.idle":"2023-10-27T08:57:10.951377Z","shell.execute_reply.started":"2023-10-27T08:57:10.892804Z","shell.execute_reply":"2023-10-27T08:57:10.950608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🏃 Run Training\n每一轮训练过后，需要做一些实验数据相关的统计，例如训练过程中的最高准确率、最高准确率对应的训练轮次，模型、每一轮训练产生的训练损失、每一轮训练过后模型在验证集上的损失大小和准确率等等，通过 $wandb.log$ 方法更新相关数值。","metadata":{"papermill":{"duration":0.015603,"end_time":"2023-10-27T00:44:56.055599","exception":false,"start_time":"2023-10-27T00:44:56.039996","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def run_training (model, run, optimizer, scheduler, device, num_epochs) :\n    # wandb 自动记录 PyTorch 模型的权重、偏置和梯度, log_freq 设置记录的频率 (每 log_freq 批次记录一次)\n    wandb.watch (model, log_freq = 100)\n    \n    # 打印 GPU 的名字\n    if torch.cuda.is_available () :\n        print (\"cuda: {}\\n\".format (torch.cuda.get_device_name ()))\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    for epoch in range (1, num_epochs + 1): \n        gc.collect ()\n        print (f'Epoch {epoch} / {num_epochs}', end = '')\n        train_loss = train_one_epoch (model, optimizer, scheduler, \n                                      dataloader = train_loader, \n                                      device = CFG.device, epoch = epoch)\n  \n\n        val_loss, val_accuracy = valid_one_epoch (model, valid_loader,\n                                                  device = CFG.device, \n                                                  epoch = epoch)\n        history['Train Loss'].append (train_loss)\n        history['Valid Loss'].append (val_loss)\n        history['Valid Accuracy'].append (val_accuracy)\n        \n        # 记录、更新模型的性能指标\n        wandb.log ({\"Train Loss\" : train_loss, \n                   \"Valid Loss\" : val_loss,\n                   \"Valid Accuracy\" : val_accuracy,\n                   \"LR\" : scheduler.get_last_lr ()[0]})\n        \n        # 保存更优的模型\n        if val_accuracy >= best_accuracy :\n            print(f\"{c_}Valid Accuracy Improved ({best_accuracy:0.4f} ---> {val_accuracy:0.4f})\")\n            best_accuracy = val_accuracy\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_model.bin\"\n            torch.save (model.state_dict (), PATH)\n            # Save a model file from the current directory\n            wandb.save (PATH)\n            print (f\"Model Saved{sr_}\")\n            \n        last_model_wts = copy.deepcopy (model.state_dict ())\n        PATH = f\"last_model.bin\"\n        torch.save (model.state_dict (), PATH)\n            \n        print (); print ()\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    # load best model weights\n    model.load_state_dict (best_model_wts)\n    \n    return model, history","metadata":{"papermill":{"duration":0.094917,"end_time":"2023-10-27T00:44:56.166039","exception":false,"start_time":"2023-10-27T00:44:56.071122","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-10-27T08:57:10.952646Z","iopub.execute_input":"2023-10-27T08:57:10.952942Z","iopub.status.idle":"2023-10-27T08:57:11.014152Z","shell.execute_reply.started":"2023-10-27T08:57:10.952919Z","shell.execute_reply":"2023-10-27T08:57:11.013325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🚅 Training\n真的开始训练了。","metadata":{"papermill":{"duration":0.014948,"end_time":"2023-10-27T00:44:56.259571","exception":false,"start_time":"2023-10-27T00:44:56.244623","status":"completed"},"tags":[]}},{"cell_type":"code","source":"run = wandb.init (project = CFG.exp_name, \n                 config = {k : v for k, v in dict (vars (CFG)).items () if '__' not in k},\n                 anonymous = anonymous,\n                 name = f\"model-{CFG.model_name}|dim-{CFG.img_size[0]}x{CFG.img_size[1]}|optimizer-{CFG.optimizer}\",\n                )\ntrain_loader, valid_loader, test_loader = DataLoader (data_transforms)\nmodel     = CNN (ConvBlock, DenseBlock).to (CFG.device)\noptimizer = torch.optim.Adam (model.parameters (), lr = CFG.lr, weight_decay = CFG.wd)\nscheduler = fetch_scheduler (optimizer)\nmodel, history = run_training (model, run, optimizer, scheduler,\n                              device = CFG.device,\n                              num_epochs = CFG.epochs)\nrun.finish ()\ndisplay (ipd.IFrame (run.url, width = 1000, height = 720))","metadata":{"papermill":{"duration":27474.117247,"end_time":"2023-10-27T08:22:50.392194","exception":false,"start_time":"2023-10-27T00:44:56.274947","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-10-27T08:57:11.015234Z","iopub.execute_input":"2023-10-27T08:57:11.015495Z","iopub.status.idle":"2023-10-27T09:24:41.494928Z","shell.execute_reply.started":"2023-10-27T08:57:11.015472Z","shell.execute_reply":"2023-10-27T09:24:41.493928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🔭 Prediction\n从测试集中选取若干张图片进行预测。从预测结果来看，模型的分类能力还是可观的，三十张图片中主要有一个错误，第三行第一列并没有在安全驾驶，可能在和乘客交流；第三行第三列我自己也判断不出来司机的行为，是对是错也不清楚。","metadata":{"papermill":{"duration":0.032324,"end_time":"2023-10-27T08:22:50.456730","exception":false,"start_time":"2023-10-27T08:22:50.424406","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def Precision (model, device) :\n    _, valid_loader, test_loader = DataLoader (data_transforms)\n    it = iter (test_loader)\n    for i in range (3) :\n        imgs = next (it)\n        imgs = imgs.to (device, dtype = torch.float)\n        preds = model (imgs)\n        _, preds = torch.max (preds, dim = 1)\n        showImage (imgs, preds, size = 5)\n        \nPrecision (model, device = CFG.device)","metadata":{"papermill":{"duration":0.111918,"end_time":"2023-10-27T08:22:50.601069","exception":true,"start_time":"2023-10-27T08:22:50.489151","status":"failed"},"tags":[],"execution":{"iopub.status.busy":"2023-10-27T09:24:41.496190Z","iopub.execute_input":"2023-10-27T09:24:41.496484Z","iopub.status.idle":"2023-10-27T09:24:45.255859Z","shell.execute_reply.started":"2023-10-27T09:24:41.496458Z","shell.execute_reply":"2023-10-27T09:24:45.254966Z"},"trusted":true},"execution_count":null,"outputs":[]}]}