{"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":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nfrom glob import glob\nimport numpy as np\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, random_split\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import r2_score, mean_squared_error\nimport torchvision.models as models\n\n# 配置参数\nIMG_SIZE = 224\nCOLOR_TYPE = 1  # 1表示灰度，3表示彩色\nNUM_CLASSES = 10\nBATCH_SIZE = 128\nEPOCHS = 15\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# ---------------- CBAM模块定义 ----------------\nclass ChannelAttention(nn.Module):\n    def __init__(self, in_planes, ratio=8):\n        super(ChannelAttention, self).__init__()\n        self.avg_pool = nn.AdaptiveAvgPool2d(1)\n        self.max_pool = nn.AdaptiveMaxPool2d(1)\n        self.fc = nn.Sequential(\n            nn.Conv2d(in_planes, in_planes // ratio, 1, bias=False),\n            nn.ReLU(),\n            nn.Conv2d(in_planes // ratio, in_planes, 1, bias=False)\n        )\n        self.sigmoid = nn.Sigmoid()\n    def forward(self, x):\n        avg_out = self.fc(self.avg_pool(x))\n        max_out = self.fc(self.max_pool(x))\n        out = avg_out + max_out\n        return self.sigmoid(out)\n\nclass SpatialAttention(nn.Module):\n    def __init__(self, kernel_size=7):\n        super(SpatialAttention, self).__init__()\n        self.conv1 = nn.Conv2d(2, 1, kernel_size, padding=kernel_size//2, bias=False)\n        self.sigmoid = nn.Sigmoid()\n    def forward(self, x):\n        avg_out = torch.mean(x, dim=1, keepdim=True)\n        max_out, _ = torch.max(x, dim=1, keepdim=True)\n        x_cat = torch.cat([avg_out, max_out], dim=1)\n        x_out = self.conv1(x_cat)\n        return self.sigmoid(x_out)\n\nclass CBAM(nn.Module):\n    def __init__(self, channels, ratio=8, kernel_size=7):\n        super(CBAM, self).__init__()\n        self.ca = ChannelAttention(channels, ratio)\n        self.sa = SpatialAttention(kernel_size)\n    def forward(self, x):\n        x = x * self.ca(x)\n        x = x * self.sa(x)\n        return x\n\n# ---------------- 数据集定义 ----------------\nclass DrivingDataset(Dataset):\n    def __init__(self, root_dir, mode='train'):\n        self.images = []\n        self.labels = []\n        if mode == 'train':\n            for class_id in range(NUM_CLASSES):\n                files = glob(f'{root_dir}/imgs/train/c{class_id}/*.jpg')\n                for file in files:\n                    img = cv2.imread(file, cv2.IMREAD_GRAYSCALE if COLOR_TYPE == 1 else cv2.IMREAD_COLOR)\n                    img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n                    if COLOR_TYPE == 1:\n                        img = np.expand_dims(img, axis=-1)\n                    self.images.append(img)\n                    self.labels.append(class_id)\n        else:\n            files = sorted(glob(f'{root_dir}/imgs/test/*.jpg'))\n            for file in files:\n                img = cv2.imread(file, cv2.IMREAD_GRAYSCALE if COLOR_TYPE == 1 else cv2.IMREAD_COLOR)\n                img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n                if COLOR_TYPE == 1:\n                    img = np.expand_dims(img, axis=-1)\n                self.images.append(img)\n                self.labels.append(-1)\n        self.images = np.array(self.images, dtype=np.float32) / 255.0\n        if COLOR_TYPE == 1:\n            self.images = self.images.transpose((0, 3, 1, 2))\n        else:\n            self.images = self.images.transpose((0, 3, 1, 2))\n        self.labels = np.array(self.labels, dtype=np.int64)\n\n    def __len__(self):\n        return len(self.images)\n\n    def __getitem__(self, idx):\n        return torch.tensor(self.images[idx]), torch.tensor(self.labels[idx])\n\n# ---------------- MobileNetV2+CBAM模型定义 ----------------\nclass Bottleneck_CBAM(nn.Module):\n    def __init__(self, block, cbam):\n        super().__init__()\n        self.block = block\n        self.cbam = cbam\n\n    def forward(self, x):\n        out = self.block(x)\n        out = self.cbam(out)\n        return out\n\nclass MobileNetV2_CBAM(nn.Module):\n    def __init__(self, num_classes=NUM_CLASSES, color_type=COLOR_TYPE):\n        super(MobileNetV2_CBAM, self).__init__()\n        base_model = models.mobilenet_v2(pretrained=True)\n        if color_type == 1:\n            base_model.features[0][0] = nn.Conv2d(1, 32, kernel_size=3, stride=2, padding=1, bias=False)\n        # 用CBAM包装每个bottleneck block\n        for i, m in enumerate(base_model.features):\n            if isinstance(m, nn.Sequential) or isinstance(m, nn.Conv2d):\n                continue\n            if hasattr(m, 'out_channels'):\n                out_c = m.out_channels\n            elif hasattr(m, 'conv'):\n                out_c = m.conv[-1].out_channels\n            else:\n                out_c = 32\n            base_model.features[i] = Bottleneck_CBAM(m, CBAM(out_c))\n        base_model.classifier[1] = nn.Linear(base_model.last_channel, num_classes)\n        self.model = base_model\n\n    def forward(self, x):\n        return self.model(x)\n\n# ---------------- 训练过程指标可视化 ----------------\ndef plot_metrics(train_loss, val_loss, train_acc, val_acc, train_r2, val_r2, train_mse, val_mse, save_dir='results'):\n    os.makedirs(save_dir, exist_ok=True)\n    plt.figure()\n    plt.plot(train_loss, label='Train Loss')\n    plt.plot(val_loss, label='Val Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.legend()\n    plt.title('Training and Validation Loss')\n    plt.savefig(os.path.join(save_dir, 'loss.png'))\n    plt.close()\n\n    plt.figure()\n    plt.plot(train_acc, label='Train Acc')\n    plt.plot(val_acc, label='Val Acc')\n    plt.xlabel('Epoch')\n    plt.ylabel('Accuracy')\n    plt.legend()\n    plt.title('Training and Validation Accuracy')\n    plt.savefig(os.path.join(save_dir, 'accuracy.png'))\n    plt.close()\n\n    plt.figure()\n    plt.plot(train_r2, label='Train R2')\n    plt.plot(val_r2, label='Val R2')\n    plt.xlabel('Epoch')\n    plt.ylabel('R2')\n    plt.legend()\n    plt.title('Training and Validation R2')\n    plt.savefig(os.path.join(save_dir, 'r2.png'))\n    plt.close()\n\n    plt.figure()\n    plt.plot(train_mse, label='Train MSE')\n    plt.plot(val_mse, label='Val MSE')\n    plt.xlabel('Epoch')\n    plt.ylabel('MSE')\n    plt.legend()\n    plt.title('Training and Validation MSE')\n    plt.savefig(os.path.join(save_dir, 'mse.png'))\n    plt.close()\n\n# ---------------- Grad-CAM热力图叠加显示 ----------------\ndef show_gradcam_on_image(img, mask, alpha=0.5):\n    if img.shape[0] == 1:\n        img = np.repeat(img, 3, axis=0)\n    img = np.transpose(img, (1, 2, 0))\n    img = np.uint8(255 * img)\n    heatmap = cv2.applyColorMap(np.uint8(255 * mask), cv2.COLORMAP_JET)\n    heatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB)\n    superimposed_img = cv2.addWeighted(img, 1 - alpha, heatmap, alpha, 0)\n    plt.figure(figsize=(4, 4))\n    plt.imshow(superimposed_img)\n    plt.axis('off')\n    plt.title('Grad-CAM')\n    plt.show()\n\n# ---------------- Grad-CAM生成函数 ----------------\ndef generate_gradcam(model, input_tensor, target_class, conv_layer=None):\n    model.eval()\n    input_tensor = input_tensor.unsqueeze(0).to(DEVICE)\n    # 默认用最后一层卷积层\n    if conv_layer is None:\n        last_conv = model.model.features[-1][0]\n    else:\n        last_conv = conv_layer\n    activations = []\n    gradients = []\n\n    def forward_hook(module, input, output):\n        activations.append(output.detach())\n\n    def backward_hook(module, grad_in, grad_out):\n        gradients.append(grad_out[0].detach())\n\n    handle_f = last_conv.register_forward_hook(forward_hook)\n    handle_b = last_conv.register_full_backward_hook(backward_hook)\n\n    output = model(input_tensor)\n    pred_class = output.argmax(dim=1).item()\n    class_idx = target_class if target_class is not None else pred_class\n    score = output[0, class_idx]\n    model.zero_grad()\n    score.backward()\n\n    grads_val = gradients[0]\n    activations_val = activations[0]\n    weights = grads_val.mean(dim=(2, 3), keepdim=True)\n    cam = (weights * activations_val).sum(dim=1, keepdim=True)\n    cam = torch.relu(cam)\n    cam = cam.squeeze().cpu().numpy()\n    cam = cv2.resize(cam, (IMG_SIZE, IMG_SIZE))\n    cam = (cam - cam.min()) / (cam.max() - cam.min() + 1e-8)\n\n    handle_f.remove()\n    handle_b.remove()\n    return cam\n\n# ---------------- 训练与验证主流程 ----------------\ndef train_model(model, train_loader, val_loader, epochs=10, save_dir='results'):\n    criterion = nn.CrossEntropyLoss()\n    optimizer = optim.Adam(model.parameters(), lr=1e-4)\n    train_losses, val_losses = [], []\n    train_accs, val_accs = [], []\n    train_r2s, val_r2s = [], []\n    train_mses, val_mses = [], []\n    best_val_acc = 0.0\n    os.makedirs(save_dir, exist_ok=True)\n    for epoch in range(epochs):\n        model.train()\n        running_loss = 0.0\n        correct = 0\n        total = 0\n        all_labels = []\n        all_preds = []\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n            running_loss += loss.item() * imgs.size(0)\n            preds = outputs.argmax(1)\n            correct += (preds == labels).sum().item()\n            total += labels.size(0)\n            all_labels.extend(labels.cpu().numpy())\n            all_preds.extend(preds.cpu().numpy())\n        train_loss = running_loss / len(train_loader.dataset)\n        train_acc = correct / total\n        train_losses.append(train_loss)\n        train_accs.append(train_acc)\n        train_r2s.append(r2_score(all_labels, all_preds))\n        train_mses.append(mean_squared_error(all_labels, all_preds))\n\n        model.eval()\n        val_loss = 0.0\n        correct = 0\n        total = 0\n        all_labels = []\n        all_preds = []\n        val_imgs_list = []\n        val_labels_list = []\n        with torch.no_grad():\n            for imgs, labels in val_loader:\n                imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n                outputs = model(imgs)\n                loss = criterion(outputs, labels)\n                val_loss += loss.item() * imgs.size(0)\n                preds = outputs.argmax(1)\n                correct += (preds == labels).sum().item()\n                total += labels.size(0)\n                all_labels.extend(labels.cpu().numpy())\n                all_preds.extend(preds.cpu().numpy())\n                val_imgs_list.append(imgs.cpu())\n                val_labels_list.append(labels.cpu())\n        val_loss /= len(val_loader.dataset)\n        val_acc = correct / total\n        val_losses.append(val_loss)\n        val_accs.append(val_acc)\n        val_r2s.append(r2_score(all_labels, all_preds))\n        val_mses.append(mean_squared_error(all_labels, all_preds))\n        print(f\"Epoch {epoch+1}/{epochs} | Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f} | \"\n              f\"Train Acc: {train_acc:.4f} | Val Acc: {val_acc:.4f} | \"\n              f\"Train R2: {train_r2s[-1]:.4f} | Val R2: {val_r2s[-1]:.4f} | \"\n              f\"Train MSE: {train_mses[-1]:.4f} | Val MSE: {val_mses[-1]:.4f}\")\n\n        # 保存最优模型\n        if val_acc > best_val_acc:\n            best_val_acc = val_acc\n            torch.save(model.state_dict(), os.path.join(save_dir, 'best_MobileNetV2_CBAM_model.pth'))\n            print(f\"Best model saved at epoch {epoch+1} with val_acc={val_acc:.4f}\")\n\n            # 只在保存最优模型时输出按类别准确率\n            class_correct = [0 for _ in range(NUM_CLASSES)]\n            class_total = [0 for _ in range(NUM_CLASSES)]\n            for label, pred in zip(all_labels, all_preds):\n                if 0 <= label < NUM_CLASSES:\n                    class_total[label] += 1\n                    if label == pred:\n                        class_correct[label] += 1\n            print(\"Per-class accuracy (only for best model):\")\n            for i in range(NUM_CLASSES):\n                acc = class_correct[i] / class_total[i] if class_total[i] > 0 else 0\n                print(f\"  Class {i}: {acc:.4f} ({class_correct[i]}/{class_total[i]})\")\n\n    # 随机选取10张验证集图片生成Grad-CAM热力图\n    val_imgs_all = torch.cat(val_imgs_list, dim=0)\n    val_labels_all = torch.cat(val_labels_list, dim=0)\n    idxs = np.random.choice(val_imgs_all.shape[0], 10, replace=False)\n    for i, idx in enumerate(idxs):\n        img = val_imgs_all[idx].numpy()\n        label = val_labels_all[idx].item()\n        cam = generate_gradcam(model, val_imgs_all[idx], target_class=label)\n        plt.figure(figsize=(4, 4))\n        show_gradcam_on_image(img, cam)\n        plt.savefig(os.path.join(save_dir, f'gradcam_{i+1}.png'))\n        plt.close()\n    return train_losses, val_losses, train_accs, val_accs, train_r2s, val_r2s, train_mses, val_mses\n\n# ---------------- 主函数 ----------------\ndef main():\n    root_dir = '../input/state-farm-distracted-driver-detection'\n    dataset = DrivingDataset(root_dir, mode='train')\n    train_size = int(0.8 * len(dataset))\n    val_size = len(dataset) - train_size\n    train_set, val_set = random_split(dataset, [train_size, val_size])\n    train_loader = DataLoader(train_set, batch_size=BATCH_SIZE, shuffle=True)\n    val_loader = DataLoader(val_set, batch_size=BATCH_SIZE)\n\n    print(\"Training MobileNetV2_CBAM...\")\n    model = MobileNetV2_CBAM().to(DEVICE)\n    save_dir = 'results'\n    train_losses, val_losses, train_accs, val_accs, train_r2s, val_r2s, train_mses, val_mses = train_model(\n        model, train_loader, val_loader, EPOCHS, save_dir=save_dir)\n    plot_metrics(train_losses, val_losses, train_accs, val_accs, train_r2s, val_r2s, train_mses, val_mses, save_dir=save_dir)\n\nif __name__ == \"__main__\":\n    main()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null}]}