{"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":10902273,"sourceType":"datasetVersion","datasetId":6775873}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchvision.transforms as transforms\nfrom torchvision.datasets import ImageFolder\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn.functional as F\nfrom transformers import ViTForImageClassification, ViTImageProcessor\nimport pandas as pd\nfrom sklearn.model_selection import train_test_split\nfrom PIL import Image\nfrom tqdm import tqdm\nfrom sklearn.preprocessing import LabelEncoder\nimport os\nimport shutil\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint('Device: ',device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-09T11:40:05.283001Z","iopub.execute_input":"2025-03-09T11:40:05.283296Z","iopub.status.idle":"2025-03-09T11:40:05.289105Z","shell.execute_reply.started":"2025-03-09T11:40:05.283275Z","shell.execute_reply":"2025-03-09T11:40:05.288159Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\ndf = pd.read_csv(\"/kaggle/input/split-data-ver-3/split_data/val.csv\")\nclass_counts = df[\"whaleID\"].value_counts()\nvalid_classes = class_counts[class_counts >= 2].index\ndf_filtered = df[df[\"whaleID\"].isin(valid_classes)]\ndf_filtered.to_csv(\"/kaggle/working/val.csv\", index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-09T11:39:39.813469Z","iopub.execute_input":"2025-03-09T11:39:39.813774Z","iopub.status.idle":"2025-03-09T11:39:39.825404Z","shell.execute_reply.started":"2025-03-09T11:39:39.813750Z","shell.execute_reply":"2025-03-09T11:39:39.824599Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_path1 = \"/kaggle/input/split-data-ver-3/split_data/train\"\nimage_path2 = \"/kaggle/input/split-data-ver-3/split_data/aug\"\nimage_path3 = \"/kaggle/input/split-data-ver-3/split_data/val\"\n\ncsv_paths = [\n    (\"/kaggle/input/split-data-ver-3/split_data/train_split.csv\", image_path1),\n    (\"/kaggle/input/split-data-ver-3/split_data/augmented_labels.csv\", image_path2),\n    (\"/kaggle/working/val.csv\", image_path3)\n]\n\ndf_list = []\nfor csv_path, img_path in csv_paths:\n    df = pd.read_csv(csv_path, names=[\"image_path\", \"whaleID\"])\n    df = df.iloc[1:].reset_index(drop=True)\n    df[\"base_path\"] = img_path\n    df_list.append(df)\n\ndf = pd.concat(df_list, ignore_index=True)\ndf = df.iloc[1:].reset_index(drop=True)\ndf = df[~df[\"whaleID\"].isin([1, 247, 397])]\ndf.head(10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-09T11:40:05.485470Z","iopub.execute_input":"2025-03-09T11:40:05.485797Z","iopub.status.idle":"2025-03-09T11:40:05.517975Z","shell.execute_reply.started":"2025-03-09T11:40:05.485769Z","shell.execute_reply":"2025-03-09T11:40:05.516977Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Tổng số ảnh:\", len(df))\nprint(\"Số lượng lớp unique whaleID:\", df[\"whaleID\"].nunique())\nprint(df.head(10))\nprint(\"Số lượng whaleID = -1:\", (df[\"whaleID\"] == -1).sum())\nprint(df[\"whaleID\"].value_counts().tail(10))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-09T11:40:08.357871Z","iopub.execute_input":"2025-03-09T11:40:08.358249Z","iopub.status.idle":"2025-03-09T11:40:08.372676Z","shell.execute_reply.started":"2025-03-09T11:40:08.358220Z","shell.execute_reply":"2025-03-09T11:40:08.371834Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_encoder = LabelEncoder()\ndf[\"whaleID\"] = label_encoder.fit_transform(df[\"whaleID\"])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-09T11:40:10.233575Z","iopub.execute_input":"2025-03-09T11:40:10.233859Z","iopub.status.idle":"2025-03-09T11:40:10.241513Z","shell.execute_reply.started":"2025-03-09T11:40:10.233838Z","shell.execute_reply":"2025-03-09T11:40:10.240597Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_counts = df[\"whaleID\"].value_counts()\nvalid_classes = class_counts[class_counts >= 2].index\ndf = df[df[\"whaleID\"].isin(valid_classes)]\n\nprint(class_counts)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-09T11:40:11.761537Z","iopub.execute_input":"2025-03-09T11:40:11.761829Z","iopub.status.idle":"2025-03-09T11:40:11.769974Z","shell.execute_reply.started":"2025-03-09T11:40:11.761808Z","shell.execute_reply":"2025-03-09T11:40:11.769193Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"feature_extractor = ViTImageProcessor.from_pretrained(\"google/vit-base-patch16-224-in21k\")\ntrain_df, temp_df = train_test_split(df, test_size=0.3, random_state=42, stratify=df[\"whaleID\"])\nval_df, test_df = train_test_split(temp_df, test_size=1/3, random_state=42, stratify=temp_df[\"whaleID\"])\n\nprint(f\"Train: {len(train_df)}, Val: {len(val_df)}, Test: {len(test_df)}\")\nprint(train_df.head())\nprint(val_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-09T11:40:14.282349Z","iopub.execute_input":"2025-03-09T11:40:14.282643Z","iopub.status.idle":"2025-03-09T11:40:14.567919Z","shell.execute_reply.started":"2025-03-09T11:40:14.282621Z","shell.execute_reply":"2025-03-09T11:40:14.567189Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class WhaleDataset(Dataset):\n    def __init__(self, df):\n        self.df = df\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        img_path = os.path.join(self.df.iloc[idx]['base_path'], self.df.iloc[idx]['image_path'])\n        label = self.df.iloc[idx]['whaleID']\n\n        if not os.path.exists(img_path):\n            print(f\"⚠️ Ảnh không tồn tại: {img_path}\")\n            return None \n\n        image = Image.open(img_path).convert(\"RGB\")\n        image = feature_extractor(image, return_tensors=\"pt\")[\"pixel_values\"].squeeze(0)\n\n        return image, torch.tensor(label, dtype=torch.long)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-09T11:40:18.378708Z","iopub.execute_input":"2025-03-09T11:40:18.379022Z","iopub.status.idle":"2025-03-09T11:40:18.384515Z","shell.execute_reply.started":"2025-03-09T11:40:18.378999Z","shell.execute_reply":"2025-03-09T11:40:18.383577Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = WhaleDataset(train_df)\nval_dataset = WhaleDataset(val_df)\ntest_dataset = WhaleDataset(test_df)\n\ntrain_loader = DataLoader(train_dataset, batch_size=16, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=16, shuffle=False)\ntest_loader = DataLoader(test_dataset, batch_size=16, shuffle=False)\n\nprint(f\"Tập train: {len(train_loader)} batches, Tập val: {len(val_loader)} batches, Tập test: {len(test_loader)} batches\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-09T11:40:18.581133Z","iopub.execute_input":"2025-03-09T11:40:18.581418Z","iopub.status.idle":"2025-03-09T11:40:18.587357Z","shell.execute_reply.started":"2025-03-09T11:40:18.581395Z","shell.execute_reply":"2025-03-09T11:40:18.586493Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_classes = df[\"whaleID\"].nunique()\n\nclass CustomViT(nn.Module):\n    def __init__(self, num_classes):\n        super(CustomViT, self).__init__()\n        self.vit = ViTForImageClassification.from_pretrained(\n            \"google/vit-base-patch16-224-in21k\",\n            num_labels=num_classes\n        )\n        for layer in self.vit.vit.encoder.layer:\n            layer.attention.attention.dropout = nn.Dropout(0.5)\n            layer.output.dropout = nn.Dropout(0.5)\n            \n        self.layernorm = nn.LayerNorm(self.vit.config.hidden_size)\n        self.relu = nn.ReLU()\n        self.dropout = nn.Dropout(p=0.5)\n\n        self.classifier = nn.Linear(self.vit.config.hidden_size, num_classes)\n\n    def forward(self, x):\n        outputs = self.vit.vit(x).last_hidden_state[:, 0, :]\n        outputs = self.layernorm(outputs)  \n        outputs = self.relu(outputs) \n        outputs = self.dropout(outputs) \n        logits = self.classifier(outputs)\n        return {\"logits\": logits}\n\nmodel = CustomViT(num_classes).to(device)\n\ncriterion = nn.CrossEntropyLoss()\noptimizer = optim.AdamW(model.parameters(), lr=5e-5)\n\nepochs = 30","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-09T08:29:38.878349Z","iopub.execute_input":"2025-03-09T08:29:38.878670Z","iopub.status.idle":"2025-03-09T08:29:41.426581Z","shell.execute_reply.started":"2025-03-09T08:29:38.878643Z","shell.execute_reply":"2025-03-09T08:29:41.425836Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_losses = []\nval_losses = []\ntrain_accuracies = []\nval_accuracies = []\nbest_val_loss = float('inf')\ntrigger_times = 0\npatience = 3\n\nfor epoch in range(epochs):\n    model.train()\n    train_loss, correct = 0, 0\n\n    for images, labels in train_loader:\n        images, labels = images.to(device), labels.to(device)\n\n        optimizer.zero_grad()\n        outputs = model(images)[\"logits\"]\n        loss = criterion(outputs, labels)\n\n        loss.backward()\n        optimizer.step()\n\n        train_loss += loss.item()\n        correct += (outputs.argmax(1) == labels).sum().item()\n\n    train_acc = correct / len(train_loader.dataset)\n    train_losses.append(train_loss / len(train_loader))\n    train_accuracies.append(train_acc)\n\n    model.eval()\n    val_loss, val_correct = 0, 0\n\n    with torch.no_grad():\n        for images, labels in val_loader:\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)[\"logits\"]\n            loss = criterion(outputs, labels)\n\n            val_loss += loss.item()\n            val_correct += (outputs.argmax(1) == labels).sum().item()\n\n    val_acc = val_correct / len(val_loader.dataset)\n    val_losses.append(val_loss / len(val_loader))\n    val_accuracies.append(val_acc)\n\n    print(f\"✅ Epoch {epoch+1}/{epochs} - Train Loss: {train_losses[-1]:.4f}, Train Acc: {train_accuracies[-1]:.4f}, Val Loss: {val_losses[-1]:.4f}, Val Acc: {val_accuracies[-1]:.4f}\")\n\n    if val_losses[-1] < best_val_loss:\n        best_val_loss = val_losses[-1]\n        trigger_times = 0\n    else:\n        trigger_times += 1\n\n    if trigger_times >= patience:\n        print(f\"🛑 Early stopping: validation loss không được cải thiện sau {patience} epoch liên tục!\")\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-08T00:55:22.352004Z","iopub.execute_input":"2025-03-08T00:55:22.352217Z","execution_failed":"2025-03-08T01:23:19.072Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_path = '/kaggle/working/model_whale_ver1.pth'\ntorch.save(model.state_dict(), model_path)\n\nepochs_range = range(1, len(train_accuracies) + 1)\n\nplt.figure(figsize=(12, 6))\nsns.set_style(\"darkgrid\")\n\nplt.subplot(1, 2, 1)\nplt.plot(epochs_range, train_accuracies, label=\"Train Accuracy\", marker='o')\nplt.plot(epochs_range, val_accuracies, label=\"Val Accuracy\", marker='o')\nplt.xlabel(\"Epochs\")\nplt.ylabel(\"Loss\")\nplt.title(\"Training & Validation Accuracy\")\nplt.legend()\n\nplt.subplot(1, 2, 2)\nplt.plot(epochs_range, train_losses, label=\"Train Loss\", marker='o')\nplt.plot(epochs_range, val_losses, label=\"Val Loss\", marker='o')\nplt.xlabel(\"Epochs\")\nplt.ylabel(\"Accuracy\")\nplt.title(\"Training & Validation Loss\")\nplt.legend()\n\nplt.savefig(\"/kaggle/working/training_plot.jpg\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-09T08:29:45.158500Z","iopub.execute_input":"2025-03-09T08:29:45.158802Z","iopub.status.idle":"2025-03-09T08:29:45.901128Z","shell.execute_reply.started":"2025-03-09T08:29:45.158780Z","shell.execute_reply":"2025-03-09T08:29:45.900053Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_path = '/kaggle/working/model_whale_ver1.pth'\n\nmodel = CustomViT(num_classes)\nmodel.load_state_dict(torch.load(model_path, map_location=device))\nmodel.to(device)\nmodel.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-09T08:30:00.305171Z","iopub.execute_input":"2025-03-09T08:30:00.305486Z","iopub.status.idle":"2025-03-09T08:30:01.370569Z","shell.execute_reply.started":"2025-03-09T08:30:00.305459Z","shell.execute_reply":"2025-03-09T08:30:01.369818Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"correct = 0\ntotal = 0\n\npredictions = []\ntrue_labels = []\nimages_list = []\n\nwith torch.no_grad():\n    for images, labels in test_loader:\n        images, labels = images.to(device), labels.to(device)\n\n        outputs = model(images)[\"logits\"]\n        _, predicted = torch.max(outputs, 1)\n\n        correct += (predicted == labels).sum().item()\n        total += labels.size(0)\n\n        # Lưu 10 ảnh đầu tiên để visualize\n        for i in range(min(10, images.shape[0])):\n            images_list.append(images[i].cpu().numpy())\n            predictions.append(predicted[i].item())\n            true_labels.append(labels[i].item())\n\ntest_accuracy = correct / total * 100\nprint(f\"🎯 Test Accuracy: {test_accuracy:.2f}%\")\n\nfig, axes = plt.subplots(2, 5, figsize=(15, 6))\n\nfor i in range(10):\n    img = images_list[i].transpose(1, 2, 0)\n    img = (img - img.min()) / (img.max() - img.min())\n    ax = axes[i // 5, i % 5]\n    ax.imshow(img)\n    ax.axis(\"off\")\n    ax.set_title(f\"Pred: {predictions[i]}\\nTrue: {true_labels[i]}\", \n                 color=\"green\" if predictions[i] == true_labels[i] else \"red\")\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-09T08:30:17.255871Z","iopub.execute_input":"2025-03-09T08:30:17.256185Z","iopub.status.idle":"2025-03-09T08:30:28.525864Z","shell.execute_reply.started":"2025-03-09T08:30:17.256163Z","shell.execute_reply":"2025-03-09T08:30:28.524997Z"}},"outputs":[],"execution_count":null}]}