{"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":"none","dataSources":[{"sourceId":11848,"databundleVersionId":862157,"sourceType":"competition"}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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 matplotlib.pyplot as plt\nfrom plotly.subplots import make_subplots\nimport plotly.graph_objs as go\nimport copy\nimport os\nimport torch\nfrom PIL import Image\nfrom PIL import Image, ImageDraw\nfrom torch.utils.data import Dataset\nimport torchvision.transforms as transforms\nfrom torch.utils.data import random_split\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport torch.nn as nn\nfrom torchvision import utils\nimport pandas as pd","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":false,"execution":{"iopub.status.busy":"2025-03-25T15:24:26.628711Z","iopub.execute_input":"2025-03-25T15:24:26.629029Z","iopub.status.idle":"2025-03-25T15:24:35.716496Z","shell.execute_reply.started":"2025-03-25T15:24:26.629003Z","shell.execute_reply":"2025-03-25T15:24:35.715458Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# library which allows us to view model summary like keras/tf\n!pip install torchsummary","metadata":{"execution":{"iopub.status.busy":"2025-03-25T15:24:35.717595Z","iopub.execute_input":"2025-03-25T15:24:35.718081Z","iopub.status.idle":"2025-03-25T15:24:41.689501Z","shell.execute_reply.started":"2025-03-25T15:24:35.718050Z","shell.execute_reply":"2025-03-25T15:24:41.688075Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"labels_df = pd.read_csv('/kaggle/input/histopathologic-cancer-detection/train_labels.csv')\nprint(labels_df.head().to_markdown())","metadata":{"execution":{"iopub.status.busy":"2025-03-25T15:24:41.690631Z","iopub.execute_input":"2025-03-25T15:24:41.691035Z","iopub.status.idle":"2025-03-25T15:24:42.320763Z","shell.execute_reply.started":"2025-03-25T15:24:41.690994Z","shell.execute_reply":"2025-03-25T15:24:42.319305Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# No duplicate ids found\nlabels_df[labels_df.duplicated(keep=False)]","metadata":{"execution":{"iopub.status.busy":"2025-03-25T15:24:42.321621Z","iopub.execute_input":"2025-03-25T15:24:42.321972Z","iopub.status.idle":"2025-03-25T15:24:42.512907Z","shell.execute_reply.started":"2025-03-25T15:24:42.321937Z","shell.execute_reply":"2025-03-25T15:24:42.511509Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"labels_df['label'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2025-03-25T15:24:42.514506Z","iopub.execute_input":"2025-03-25T15:24:42.514910Z","iopub.status.idle":"2025-03-25T15:24:42.532243Z","shell.execute_reply.started":"2025-03-25T15:24:42.514874Z","shell.execute_reply":"2025-03-25T15:24:42.530949Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# define transformation that converts a PIL image into PyTorch tensors\nimport torchvision.transforms as transforms\ndata_transformer = transforms.Compose([transforms.ToTensor(),\n                                       transforms.Resize((46,46))])","metadata":{"execution":{"iopub.status.busy":"2025-03-25T15:24:42.538071Z","iopub.execute_input":"2025-03-25T15:24:42.538617Z","iopub.status.idle":"2025-03-25T15:24:42.552927Z","shell.execute_reply.started":"2025-03-25T15:24:42.538571Z","shell.execute_reply":"2025-03-25T15:24:42.551509Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.manual_seed(0) # fix random seed\n\nclass ImageDataset(Dataset):\n    def __init__(self, labels_file, img_dir, transform=None, max_get=4000):\n        data = pd.read_csv(labels_file)\n        print(f\"-> Data pd read_csv: {data.head(5)}\")\n        \n        self.data = data.head(max(4000, max_get))\n        \n        print(f\"-> Img dir: {img_dir}\")\n        self.img_dir = img_dir\n\n        self.transform = transform\n    \n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        img_path = os.path.join(self.img_dir, self.data.iloc[idx, 0])+\".tif\"\n        image = Image.open(img_path)\n        label = self.data.iloc[idx, 1] \n        # print(f\"-> getitem: {str(img_path)} : {label}\")\n        if self.transform:\n            image = self.transform(image)\n        return image, label\n\n# Define an object of the custom dataset for the train folder.\ndata_dir = '/kaggle/input/histopathologic-cancer-detection/'\ntrain_path = str(data_dir + \"train\")\nlabel_train_path = str(data_dir + \"train_labels.csv\")\nimg_dataset = ImageDataset(data_dir + \"train_labels.csv\", train_path, data_transformer, 5000)","metadata":{"execution":{"iopub.status.busy":"2025-03-25T15:24:42.555354Z","iopub.execute_input":"2025-03-25T15:24:42.555829Z","iopub.status.idle":"2025-03-25T15:24:43.293476Z","shell.execute_reply.started":"2025-03-25T15:24:42.555792Z","shell.execute_reply":"2025-03-25T15:24:43.291969Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len_img=len(img_dataset)\nlen_train=int(0.8*len_img)\nlen_val=len_img-len_train\n\n# Split Pytorch tensor\ntrain_ts,val_ts=random_split(img_dataset,\n                             [len_train,len_val]) # random split 80/20\n\nprint(\"train dataset size:\", len(train_ts))\nprint(\"validation dataset size:\", len(val_ts))","metadata":{"execution":{"iopub.status.busy":"2025-03-25T15:24:43.294801Z","iopub.execute_input":"2025-03-25T15:24:43.295265Z","iopub.status.idle":"2025-03-25T15:24:43.335095Z","shell.execute_reply.started":"2025-03-25T15:24:43.295220Z","shell.execute_reply":"2025-03-25T15:24:43.333652Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define the following transformations for the training dataset\ntr_transf = transforms.Compose([\n    transforms.RandomHorizontalFlip(p=0.5), \n    transforms.RandomVerticalFlip(p=0.5),  \n    transforms.RandomRotation(45),         \n    transforms.ToTensor()])\ntr_transf","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2025-03-25T15:24:43.336551Z","iopub.execute_input":"2025-03-25T15:24:43.337019Z","iopub.status.idle":"2025-03-25T15:24:43.391484Z","shell.execute_reply.started":"2025-03-25T15:24:43.336979Z","shell.execute_reply":"2025-03-25T15:24:43.389599Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# For the validation dataset, we don't need any augmentation; simply convert images into tensors\nval_transf = transforms.Compose([\n    transforms.ToTensor()])\n\n# After defining the transformations, overwrite the transform functions of train_ts, val_ts\ntrain_ts.transform=tr_transf\nval_ts.transform=val_transf\n\ntrain_ts, val_ts","metadata":{"execution":{"iopub.status.busy":"2025-03-25T15:24:43.393190Z","iopub.execute_input":"2025-03-25T15:24:43.393673Z","iopub.status.idle":"2025-03-25T15:24:43.412123Z","shell.execute_reply.started":"2025-03-25T15:24:43.393627Z","shell.execute_reply":"2025-03-25T15:24:43.410663Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\n# Training DataLoader\ntrain_dl = DataLoader(train_ts, \n                      batch_size=32, \n                      shuffle=True)\n\n# Validation DataLoader\nval_dl = DataLoader(val_ts,\n                    batch_size=32,\n                    shuffle=False)\n\ntrain_dl, val_dl","metadata":{"execution":{"iopub.status.busy":"2025-03-25T15:24:43.413545Z","iopub.execute_input":"2025-03-25T15:24:43.413844Z","iopub.status.idle":"2025-03-25T15:24:43.439058Z","shell.execute_reply.started":"2025-03-25T15:24:43.413820Z","shell.execute_reply":"2025-03-25T15:24:43.437515Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CNNModel.py\nimport torch.nn.functional as F\n\nclass Network(nn.Module):\n    \n    # Network Initialisation\n    def __init__(self, num_fc1 = 256, num_classes = 2, dropout_rate = 0.3):\n        \n        super(Network, self).__init__()\n    \n        self.dropout_rate= dropout_rate\n        \n        # transform resize 46 46, ảnh có chiều dài và chiều rộng là 46 x 46\n        \n        # Convolution Layers - Tầng tích chập\n        # 3, 8, 3\n        self.conv1 = nn.Conv2d(3, 16, kernel_size=3)\n        # (46 - 3 + 2*0) / 1 + 1 = 44\n        \n        # 8, 2*8, 3\n        self.conv2 = nn.Conv2d(16, 32, kernel_size=3)\n        \n        # 2*8, 4*8, 3\n        self.conv3 = nn.Conv2d(32, 64, kernel_size=3)\n        \n        # 4*8, 8*8, 3\n        self.conv4 = nn.Conv2d(64, 128, kernel_size=3)\n        \n        # Pooling layer\n        # 22 -> 10 -> 4 -> 1 => 1x1\n        self.pool = nn.MaxPool2d(kernel_size=2, stride=2) \n        \n        # Tầng fully connected\n        # VD: từ 64 kênh ở conv cuối * 1x1 conv cuối sau pooling (conv4 => pool) (64)\n        self.fc1 = nn.Linear(128 * 1 * 1, num_fc1)\n        # => 0 hoặc 1 => num_classes = 2\n        self.fc2 = nn.Linear(num_fc1, num_classes)\n\n    def forward(self,X):\n        \n        # Convolution => ReLu => Pool Layers\n        # Convolution và Relu 1: (46 - 3 + 2*0) / 1 + 1 = 44 => Pool: 44/2 = 22\n        X = self.pool(F.relu(self.conv1(X)))\n        # Convolution và Relu 2: (22 - 3 + 2*0) / 1 + 1 = 20 => Pool: 20/2 = 10\n        X = self.pool(F.relu(self.conv2(X))) \n        # Convolution và Relu 3: (10 - 3 + 2*0) / 1 + 1 = 8 => Pool: 44/2 = 4\n        X = self.pool(F.relu(self.conv3(X)))\n        # Convolution và Relu 4: (4 - 3 + 2*0) / 1 + 1 = 2 => Pool: 2/2 = 1 => 1x1\n        X = self.pool(F.relu(self.conv4(X)))\n\n        X = X.view(-1, 128 * 1 * 1)\n        \n        X = F.relu(self.fc1(X))\n        X=F.dropout(X, self.dropout_rate)\n        X = self.fc2(X)\n        return X","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:24:43.440589Z","iopub.execute_input":"2025-03-25T15:24:43.441016Z","iopub.status.idle":"2025-03-25T15:24:43.459269Z","shell.execute_reply.started":"2025-03-25T15:24:43.440975Z","shell.execute_reply":"2025-03-25T15:24:43.457910Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create instantiation of Network class\ncnn_model = Network(512, 2, 0.3)\n\n# define computation hardware approach (GPU/CPU)\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel = cnn_model.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:24:43.460857Z","iopub.execute_input":"2025-03-25T15:24:43.461335Z","iopub.status.idle":"2025-03-25T15:24:43.515133Z","shell.execute_reply.started":"2025-03-25T15:24:43.461303Z","shell.execute_reply":"2025-03-25T15:24:43.513936Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm.notebook import trange, tqdm\nfrom torch import optim\n\ndef train_model(model, train_loader, criterion, optimizer, epochs=50):\n    loss_history, accuracy_history = [], []\n    \n    for epoch in tqdm(range(epochs)):\n        model.train()\n        total_loss, correct, total = 0.0, 0, 0\n        \n        for images, labels in train_loader:\n            images, labels = images.to(device), labels.to(device)\n            optimizer.zero_grad()\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n            \n            total_loss += loss.item()\n            correct += (outputs.argmax(1) == labels).sum().item()\n            total += labels.size(0)\n        \n        loss_history.append(total_loss / len(train_loader))\n        accuracy_history.append(correct / total)\n        print(f\"Epoch {epoch+1}, Loss: {loss_history[-1]:.4f}, Accuracy: {accuracy_history[-1]:.4f}\")\n    \n    return model, loss_history, accuracy_history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:24:43.516201Z","iopub.execute_input":"2025-03-25T15:24:43.516610Z","iopub.status.idle":"2025-03-25T15:24:43.703109Z","shell.execute_reply.started":"2025-03-25T15:24:43.516580Z","shell.execute_reply":"2025-03-25T15:24:43.701949Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchsummary import summary\nsummary(cnn_model, input_size=(3, 46, 46),device=device.type)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:24:43.704510Z","iopub.execute_input":"2025-03-25T15:24:43.704913Z","iopub.status.idle":"2025-03-25T15:24:43.892983Z","shell.execute_reply.started":"2025-03-25T15:24:43.704870Z","shell.execute_reply":"2025-03-25T15:24:43.891707Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\nepochs = 50\n# cnn_model,loss_hist,metric_hist=train_model(cnn_model,train_dl,criterion,optim.Adam(cnn_model.parameters(),lr=0.005),epochs)\n# cnn_model,loss_hist,metric_hist=train_model(cnn_model,train_dl,criterion,optim.Adam(cnn_model.parameters(),lr=0.0025),epochs)\ncnn_model,loss_hist,metric_hist=train_model(cnn_model,train_dl,criterion,optim.Adam(cnn_model.parameters(),lr=0.001),epochs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:24:43.894310Z","iopub.execute_input":"2025-03-25T15:24:43.894733Z","iopub.status.idle":"2025-03-25T15:35:56.491705Z","shell.execute_reply.started":"2025-03-25T15:24:43.894696Z","shell.execute_reply":"2025-03-25T15:35:56.490243Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import seaborn as sns; sns.set(style='whitegrid')\n\nfig,ax = plt.subplots(1,2,figsize=(12,5))\n\nsns.lineplot(x=[*range(1,epochs+1)],y=loss_hist,ax=ax[0],label='loss_hist_train')\nsns.lineplot(x=[*range(1,epochs+1)],y=metric_hist,ax=ax[1],label='metric_hist_train')\n\nax[0].set_title(\"Loss Curve\")\nax[1].set_title(\"Accuracy Curve\")\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:35:56.493221Z","iopub.execute_input":"2025-03-25T15:35:56.493668Z","iopub.status.idle":"2025-03-25T15:35:57.800973Z","shell.execute_reply.started":"2025-03-25T15:35:56.493618Z","shell.execute_reply":"2025-03-25T15:35:57.799500Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate_model(model, test_loader, criterion, epochs=50):\n    model.eval()\n    loss_history, accuracy_history = [], []\n    \n    for i in tqdm(range(epochs)):\n        total_loss, correct, total = 0.0, 0, 0\n        with torch.no_grad():\n            for images, labels in test_loader:\n                images, labels = images.to(device), labels.to(device)\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n                correct += (outputs.argmax(1) == labels).sum().item()\n                total += labels.size(0)\n                total_loss += loss.item()\n        loss_history.append(total_loss / len(test_loader))\n        accuracy_history.append(correct / total)\n    \n    print(f\"Test Loss: {sum(loss_history)/len(loss_history):.4f}, Accuracy: {sum(accuracy_history)/len(accuracy_history):.4f}\")\n    return loss_history, accuracy_history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:35:57.802115Z","iopub.execute_input":"2025-03-25T15:35:57.802663Z","iopub.status.idle":"2025-03-25T15:35:57.811053Z","shell.execute_reply.started":"2025-03-25T15:35:57.802633Z","shell.execute_reply":"2025-03-25T15:35:57.809293Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"loss_val, accuracy_val = evaluate_model(cnn_model,val_dl,criterion)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:35:57.812171Z","iopub.execute_input":"2025-03-25T15:35:57.812517Z","iopub.status.idle":"2025-03-25T15:38:32.049334Z","shell.execute_reply.started":"2025-03-25T15:35:57.812487Z","shell.execute_reply":"2025-03-25T15:38:32.047858Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig,ax = plt.subplots(1,2,figsize=(12,5))\n\nsns.lineplot(x=[*range(1,epochs+1)],y=loss_val,ax=ax[0],label='loss_val')\nsns.lineplot(x=[*range(1,epochs+1)],y=accuracy_val,ax=ax[1],label='accuracy_val')\nax[0].set_title(\"Loss Val\")\nax[1].set_title(\"Accuracy Val\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:52:04.407138Z","iopub.execute_input":"2025-03-25T15:52:04.407567Z","iopub.status.idle":"2025-03-25T15:52:04.989480Z","shell.execute_reply.started":"2025-03-25T15:52:04.407535Z","shell.execute_reply":"2025-03-25T15:52:04.988183Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"loss_val, accuracy_val = evaluate_model(cnn_model,train_dl,criterion, epochs)\nfig,ax = plt.subplots(1,2,figsize=(12,5))\n\nsns.lineplot(x=[*range(1,epochs+1)],y=loss_val,ax=ax[0],label='loss_train_test')\nsns.lineplot(x=[*range(1,epochs+1)],y=accuracy_val,ax=ax[1],label='accuracy_train_test')\nax[0].set_title(\"Loss Train Test\")\nax[1].set_title(\"Accuracy Train Test\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:38:32.054015Z","iopub.execute_input":"2025-03-25T15:38:32.054329Z","iopub.status.idle":"2025-03-25T15:48:12.657123Z","shell.execute_reply.started":"2025-03-25T15:38:32.054303Z","shell.execute_reply":"2025-03-25T15:48:12.655865Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix\ndef plot_confusion_matrix(model, dataloader, device, classes):\n    model.eval()\n    all_preds = []\n    all_labels = []\n    \n    with torch.no_grad():\n        for images, labels in dataloader:   \n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            _, predicted = torch.max(outputs, 1)\n            all_preds.extend(predicted.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n    \n    cm = confusion_matrix(all_labels, all_preds)\n    plt.figure(figsize=(10, 8))\n    sns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\", xticklabels=classes, yticklabels=classes)\n    plt.xlabel('Predicted')\n    plt.ylabel('Actual')\n    plt.title('Confusion Matrix')\n    plt.show()\n\nclass_labels = [\"Không Ung Thư\", \"Ung Thư\"]\nplot_confusion_matrix(model, val_dl, device, class_labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:48:12.676653Z","iopub.execute_input":"2025-03-25T15:48:12.677086Z","iopub.status.idle":"2025-03-25T15:48:17.420432Z","shell.execute_reply.started":"2025-03-25T15:48:12.677046Z","shell.execute_reply":"2025-03-25T15:48:17.419233Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load ảnh và xử lý\npath = '/kaggle/input/histopathologic-cancer-detection/train/'\nimage_name = '5622f473549868709943855f3ee4ce5fe8a0bb4e'\nimage_path = path + image_name + \".tif\"  # Đường dẫn ảnh test\nimage = Image.open(image_path).convert(\"RGB\")  # Mở ảnh và chuyển sang RGB\nimage = data_transformer(image).unsqueeze(0)  # Chuyển thành batch 1 ảnh\n\nprint(image_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:56:43.924423Z","iopub.execute_input":"2025-03-25T15:56:43.925055Z","iopub.status.idle":"2025-03-25T15:56:43.943921Z","shell.execute_reply.started":"2025-03-25T15:56:43.925009Z","shell.execute_reply":"2025-03-25T15:56:43.942096Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Dự đoán\nwith torch.no_grad():\n    output = model(image)\n\n# Chuyển output thành lớp dự đoán\npredicted_class = torch.argmax(output, dim=1).item()\n\nprint(f\"Ảnh này thuộc lớp: {predicted_class}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:56:46.745776Z","iopub.execute_input":"2025-03-25T15:56:46.746184Z","iopub.status.idle":"2025-03-25T15:56:46.755978Z","shell.execute_reply.started":"2025-03-25T15:56:46.746156Z","shell.execute_reply":"2025-03-25T15:56:46.754596Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_prediction(image_path):  \n    image = Image.open(image_path).convert(\"RGB\")  # Mở ảnh và chuyển sang RGB\n    transformed_image = data_transformer(image).unsqueeze(0).to(device)  # Chuyển thành batch 1 ảnh và đưa vào GPU/CPU\n    \n    with torch.no_grad():\n        output = model(transformed_image)\n        predicted_class = torch.argmax(output, dim=1).item()\n    \n    # Vẽ ảnh và hiển thị dự đoán\n    plt.figure(figsize=(4, 4))\n    plt.imshow(image)  # Hiển thị ảnh gốc\n    plt.title(f\"Dự đoán: {predicted_class}\")\n    plt.axis('off')\n    plt.show()\n\n# Gọi hàm để hiển thị dự đoán cho ảnh hiện tại\nvisualize_prediction(image_path)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T15:56:48.389297Z","iopub.execute_input":"2025-03-25T15:56:48.389758Z","iopub.status.idle":"2025-03-25T15:56:48.551309Z","shell.execute_reply.started":"2025-03-25T15:56:48.389724Z","shell.execute_reply":"2025-03-25T15:56:48.549830Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# weight_path = \"weights.pt\"\n# torch.save(model.state_dict(), weight_path)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}