{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":24800,"datasetId":1042002,"databundleVersionId":1831594},{"sourceType":"datasetVersion","sourceId":14003837,"datasetId":8922549,"databundleVersionId":14780088},{"sourceType":"modelInstanceVersion","sourceId":679201,"databundleVersionId":14876975,"modelInstanceId":515210},{"sourceType":"modelInstanceVersion","sourceId":671985,"databundleVersionId":14758996,"modelInstanceId":509082},{"sourceType":"modelInstanceVersion","sourceId":678041,"databundleVersionId":14857510,"modelInstanceId":514215},{"sourceType":"modelInstanceVersion","sourceId":672244,"databundleVersionId":14762594,"modelInstanceId":509293},{"sourceType":"modelInstanceVersion","sourceId":678995,"databundleVersionId":14873394,"modelInstanceId":515036},{"sourceType":"modelInstanceVersion","sourceId":679104,"databundleVersionId":14875131,"modelInstanceId":515133},{"sourceType":"modelInstanceVersion","sourceId":673153,"databundleVersionId":14781797,"modelInstanceId":510101},{"sourceType":"modelInstanceVersion","sourceId":679218,"databundleVersionId":14877241,"modelInstanceId":515224},{"sourceType":"kernelVersion","sourceId":282312746},{"sourceType":"kernelVersion","sourceId":285211090}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Test Dataloader","metadata":{}},{"cell_type":"code","source":"# import pandas as pd\n# import numpy as np\n# import torch\n# from torch.utils.data import Dataset\n# import os\n# import cv2  # Thay pydicom bằng opencv\n# from PIL import Image\n\n\n# class VINDRCXRDataset(Dataset):\n#     def __init__(self, csv_file, image_dir, num_classes, target_size, map_size):\n#         \"\"\"\n#         Khởi tạo Dataset cho file PNG.\n#         \"\"\"\n#         # Đọc CSV\n#         self.data_df = pd.read_csv(csv_file)\n#         self.image_dir = image_dir\n#         self.target_size = target_size\n#         self.map_size = map_size\n\n#         # 1. Chuẩn bị mapping Class ID\n#         self.image_ids = self.data_df[\"image_id\"].unique()  .tolist()\n#         class_mapping_data = self.data_df[[\"class_name\", \"class_id\"]].drop_duplicates()\n\n#         self.class_to_id = {}\n#         pathology_class_count = 0\n#         for _, row in class_mapping_data.iterrows():\n#             name = row[\"class_name\"]\n#             original_id = row[\"class_id\"]\n#             # Chỉ lấy 14 bệnh lý (ID 0-13)\n#             if name.lower() != \"no finding\" and 0 <= original_id <= 13:\n#                 self.class_to_id[name] = int(original_id)\n#                 pathology_class_count += 1\n\n#         self.num_classes = pathology_class_count\n\n#         # 2. Tối ưu hóa việc tra cứu kích thước gốc (Nếu có trong CSV)\n#         # Tạo dictionary để tra cứu nhanh width/height nếu CSV có cột này\n#         self.has_dim_info = (\n#             \"width\" in self.data_df.columns and \"height\" in self.data_df.columns\n#         )\n\n#         if self.has_dim_info:\n#             # BƯỚC QUAN TRỌNG:\n#             # 1. Chỉ lấy 3 cột cần thiết\n#             # 2. Loại bỏ các dòng trùng image_id (để mỗi ảnh chỉ còn 1 dòng duy nhất chứa width/height)\n#             unique_dims = self.data_df[[\"image_id\", \"width\", \"height\"]].drop_duplicates(\n#                 subset=[\"image_id\"]\n#             )\n\n#             # 3. Giờ thì set_index sẽ an toàn và to_dict sẽ nhanh hơn nhiều\n#             self.dim_lookup = unique_dims.set_index(\"image_id\").to_dict(\"index\")\n#         else:\n#             self.dim_lookup = {}\n\n#     def __len__(self):\n#         return len(self.image_ids)\n\n#     def __getitem__(self, idx):\n#         image_id = self.image_ids[idx]\n\n#         # Lấy các dòng annotation của ảnh này\n#         img_annotations = self.data_df[self.data_df[\"image_id\"] == image_id]\n\n#         # --- 1. Load Ảnh PNG (Thay thế phần DICOM) ---\n#         img_path = os.path.join(self.image_dir, f\"{image_id}.png\")\n\n#         # Đọc ảnh grayscale (để khớp với logic cũ là 1 channel)\n#         # Nếu muốn dùng 3 kênh màu RGB thì bỏ flag cv2.IMREAD_GRAYSCALE\n#         image = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n\n#         if image is None:  # Xử lý lỗi nếu không đọc được ảnh\n#             return None, None, None, image_id\n\n#         # --- 2. Xác định Kích thước Gốc (QUAN TRỌNG ĐỂ TÍNH BBOX) ---\n#         # Nếu CSV có thông tin width/height gốc, ta ưu tiên dùng nó\n#         # (Vì nếu ảnh PNG đã bị resize trước đó, image.shape sẽ sai lệch với bbox gốc)\n#         if self.has_dim_info:\n#             dims = self.dim_lookup[image_id]\n#             original_w = dims[\"width\"]\n#             original_h = dims[\"height\"]\n#         else:\n#             # Nếu CSV không có, đành phải dùng kích thước thật của ảnh vừa đọc\n#             # Lưu ý: Nếu ảnh này đã bị resize (vd 512x512) thì bbox sẽ bị tính sai!\n#             original_h, original_w = image.shape[:2]\n\n#         if original_w <= 0 or original_h <= 0:\n#             return None, None, None, image_id\n\n#         # --- 3. Resize & Normalize ---\n#         # Resize về target_size (ví dụ 384x384)\n#         image = cv2.resize(image, (self.target_size, self.target_size))\n\n#         # Normalize 0-255 -> 0.0-1.0\n#         image = image.astype(np.float32) / 255.0\n\n#         # Chuyển sang Tensor [1, H, W]\n#         image = torch.from_numpy(image).unsqueeze(0)\n\n#         # --- 4. Tạo Label & Attention Map (Logic giữ nguyên) ---\n#         cls_label = torch.zeros(self.num_classes, dtype=torch.float32)\n#         present_classes = img_annotations[\"class_name\"].unique()\n\n#         for cls_name in present_classes:\n#             cls_id = self.class_to_id.get(cls_name)\n#             if cls_id is not None:\n#                 cls_label[cls_id] = 1.0\n\n#         # Tạo Map\n#         attn_maps = torch.zeros(\n#             (self.num_classes, self.map_size, self.map_size), dtype=torch.float32\n#         )\n\n#         for _, row in img_annotations.iterrows():\n#             cls_name = row[\"class_name\"]\n#             cls_id = self.class_to_id.get(cls_name)\n\n#             if cls_id is None:\n#                 continue\n#             if pd.isna(row[[\"x_min\", \"y_min\", \"x_max\", \"y_max\"]]).any():\n#                 continue\n\n#             # Tính tọa độ trên Map 12x12 dựa trên tỷ lệ với kích thước gốc\n#             x_min = int(row[\"x_min\"] * (self.map_size / original_w))\n#             y_min = int(row[\"y_min\"] * (self.map_size / original_h))\n#             x_max = int(row[\"x_max\"] * (self.map_size / original_w))\n#             y_max = int(row[\"y_max\"] * (self.map_size / original_h))\n\n#             x_min = max(0, x_min)\n#             y_min = max(0, y_min)\n#             x_max = min(self.map_size, x_max)\n#             y_max = min(self.map_size, y_max)\n\n#             if x_max > x_min and y_max > y_min:\n#                 attn_maps[cls_id, y_min:y_max, x_min:x_max] = 1.0\n\n#         return image, cls_label, attn_maps, image_id\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T14:06:22.554978Z","iopub.execute_input":"2025-12-10T14:06:22.555604Z","iopub.status.idle":"2025-12-10T14:06:23.069556Z","shell.execute_reply.started":"2025-12-10T14:06:22.555578Z","shell.execute_reply":"2025-12-10T14:06:23.068768Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import matplotlib.pyplot as plt\n# import cv2\n# import numpy as np\n# import torch\n\n# def visualize_sample(image, cls_label, attn_maps, image_id, class_to_id):\n#     \"\"\"\n#     image: Tensor [1, 1, H, W]\n#     attn_maps: Tensor [num_classes, map_size, map_size]\n#     cls_label: Tensor [num_classes]\n#     \"\"\"\n\n#     # --- 1. Hiển thị ảnh gốc ---\n#     img_np = image.squeeze().numpy()\n\n#     plt.figure(figsize=(6,6))\n#     plt.imshow(img_np, cmap='gray')\n#     plt.title(f\"Image ID: {image_id}\")\n#     plt.axis('off')\n#     plt.show()\n\n#     # --- 2. Hiển thị vector multi-label ---\n#     print(\"Multi-label vector:\")\n#     for cls_name, cls_id in class_to_id.items():\n#         if cls_label[0][cls_id] == 1:\n#             print(f\"  ✔ {cls_name} (id={cls_id})\")\n\n#     # --- 3. Visualize attention maps ---\n#     num_classes = attn_maps.shape[1]  # batch=1 → 1 x C x 12 x 12\n#     map_size = attn_maps.shape[-1]\n#     H, W = img_np.shape\n\n#     for cls_name, cls_id in class_to_id.items():\n#         if cls_label[0][cls_id] == 0:\n#             continue  # chỉ hiện các lớp xuất hiện\n\n#         print(f\"\\n---- Attention map for {cls_name} (id={cls_id}) ----\")\n\n#         att_map = attn_maps[0, cls_id].numpy()\n\n#         # Resize 12x12 → ảnh size (H×W)\n#         att_big = cv2.resize(att_map, (W, H), interpolation=cv2.INTER_NEAREST)\n\n#         # Normalize để overlay đẹp\n#         att_big = (att_big - att_big.min()) / (att_big.max() + 1e-6)\n\n#         # Apply heatmap\n#         heatmap = cv2.applyColorMap((att_big * 255).astype(np.uint8), cv2.COLORMAP_JET)\n#         heatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB)\n\n#         overlay = (heatmap * 0.5 + np.stack([img_np]*3, axis=-1) * 255 * 0.5).astype(np.uint8)\n\n#         plt.figure(figsize=(6,6))\n#         plt.imshow(overlay)\n#         plt.title(f\"Attention Map: {cls_name}\")\n#         plt.axis('off')\n#         plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T14:07:19.55693Z","iopub.execute_input":"2025-12-10T14:07:19.557725Z","iopub.status.idle":"2025-12-10T14:07:19.56567Z","shell.execute_reply.started":"2025-12-10T14:07:19.557699Z","shell.execute_reply":"2025-12-10T14:07:19.564903Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# from torch.utils.data import DataLoader\n\n# dataset = VINDRCXRDataset(\n#     csv_file=\"/kaggle/input/png-trainer-9/pytorch/default/1/src/csv/train_split.csv\",\n#     image_dir=\"/kaggle/input/vindr-image-convert/train_png_384\",\n#     num_classes=14,\n#     target_size=384,\n#     map_size=12\n# )\n\n# loader = DataLoader(dataset, batch_size=1, shuffle=True)\n# image, cls_label, attn_maps, image_id = next(iter(loader))\n# print(cls_label)\n# visualize_sample(image, cls_label, attn_maps, image_id[0], dataset.class_to_id)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T14:13:16.571715Z","iopub.execute_input":"2025-12-10T14:13:16.572706Z","iopub.status.idle":"2025-12-10T14:13:16.841425Z","shell.execute_reply.started":"2025-12-10T14:13:16.572678Z","shell.execute_reply":"2025-12-10T14:13:16.840721Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !python /kaggle/input/png-trainer-28/pytorch/default/1/src/train.py \\\n#   --csv_path \"/kaggle/input/png-trainer-9/pytorch/default/1/src/csv/train_split.csv\" \\\n#   --image_path \"/kaggle/input/vindr-image-convert/train_png_384\" \\\n#   --save_path \"/kaggle/working/checkpoints\" \\\n#   --batch_size 32 \\\n#   --num_prototypes 15 \\\n#   --epochs_p1 20 --lr 3e-4\n#   # --epochs_p2 15 \\\n#   # --epochs_p3 10\n\n!python /kaggle/input/png-trainer-phase-2/pytorch/default/1/src/train_phase2.py \\\n  --csv_path \"/kaggle/input/png-trainer-9/pytorch/default/1/src/csv/train_split.csv\" \\\n  --image_path \"/kaggle/input/vindr-image-convert/train_png_384\" \\\n  --save_path \"/kaggle/working/checkpoints\" \\\n--resume_phase1 /kaggle/input/csr-model-phase-1/pytorch/default/1/csr_phase1.pth \\\n--skip_phase1 \\\n  --batch_size 32 \\\n  --num_prototypes 15 \\\n--epochs_p2 20 --lr 3e-4\n      # --epochs_p3 10","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T16:43:44.406286Z","iopub.execute_input":"2025-12-10T16:43:44.406562Z","iopub.status.idle":"2025-12-10T16:43:57.696585Z","shell.execute_reply.started":"2025-12-10T16:43:44.406542Z","shell.execute_reply":"2025-12-10T16:43:57.695854Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !python /kaggle/input/evaluate-v6/pytorch/default/1/evaluate.py --test_csv /kaggle/input/inference-model-v3/pytorch/default/1/src/csv/test_split.csv --image_path /kaggle/input/vindr-image-convert/train_png_384 --checkpoint /kaggle/input/csr-model/pytorch/default/1/csr_final_model.pth --num_prototypes 10 --img_size 384 --device cuda --threshold 0.67","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T19:27:28.485472Z","iopub.execute_input":"2025-12-09T19:27:28.486175Z","iopub.status.idle":"2025-12-09T19:27:28.602076Z","shell.execute_reply.started":"2025-12-09T19:27:28.486151Z","shell.execute_reply":"2025-12-09T19:27:28.601309Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Test returned attn weights from models","metadata":{}},{"cell_type":"code","source":"# # =============================================================================\n# # INFERENCE SCRIPT CHO KAGGLE NOTEBOOK - HOÀN CHỈNH\n# # Chạy từng cell theo thứ tự\n# # =============================================================================\n\n# # %% [markdown]\n# # ## Cell 1: Import Libraries\n\n# # %%\n# import torch\n# import torch.nn as nn\n# import torch.nn.functional as F\n# import numpy as np\n# import cv2\n# import matplotlib.pyplot as plt\n# import os\n# import timm\n\n# # %% [markdown]\n# # ## Cell 2: Cấu hình - THAY ĐỔI CÁC BIẾN Ở ĐÂY\n\n# # %%\n# # =================== CẤU HÌNH - THAY ĐỔI Ở ĐÂY ===================\n\n# # Đường dẫn ảnh cần inference\n# IMAGE_PATH = \"/kaggle/input/vindr-image-convert/train_png_384/051132a778e61a86eb147c7c6f564dfe.png\"\n\n# # Đường dẫn checkpoint model\n# CHECKPOINT_PATH = \"/kaggle/input/vin-csr-training/checkpoints/csr_phase1.pth\"\n\n# # Đường dẫn lưu kết quả (None nếu không muốn lưu)\n# SAVE_PATH = \"/kaggle/working/result.png\"\n\n# # Device\n# DEVICE = \"cuda\"  # hoặc \"cpu\"\n\n# # Model config\n# NUM_CLASSES = 14\n# NUM_PROTOTYPES = 15\n# MODEL_NAME = \"densenet121\"  # hoặc \"densenet121\", \"efficientnet_b0\", ...\n# IMG_SIZE = 384\n\n# # Số bệnh top-k hiển thị\n# TOP_K = 3\n\n# # === THRESHOLD CHO SALIENCY MAPS ===\n# CAM_THRESHOLD = 1  # Giá trị từ 0.0 - 1.0 (0 = hiện tất cả, 1 = chỉ hiện vùng cao nhất)\n# SHOW_BINARY_MASK = False  # True = hiện mask nhị phân, False = hiện heatmap gradient\n\n# # Tên các bệnh\n# CLASS_NAMES = [\n#     'Aortic enlargement',\n#     'Atelectasis',\n#     'Calcification',\n#     'Cardiomegaly',\n#     'Consolidation',\n#     'ILD',\n#     'Infiltration',\n#     'Lung Opacity',\n#     'Nodule/Mass',\n#     'Other lesion',\n#     'Pleural effusion',\n#     'Pleural thickening',\n#     'Pneumothorax',\n#     'Pulmonary fibrosis'\n# ]\n\n# # ================================================================\n\n# # %% [markdown]\n# # ## Cell 3: Định nghĩa Model (CSRModel)\n\n# # %%\n# class CSRModel(nn.Module):\n#     def __init__(self, num_classes=14, num_prototypes=5, model_name=\"resnet50\", pretrained=True):\n#         super().__init__()\n        \n#         # --- PHẦN 1: CONCEPT MODEL (Giai đoạn 1) ---\n#         # F: Feature Extractor\n#         self.backbone = timm.create_model(\n#             model_name, pretrained=pretrained, features_only=True, out_indices=(4,)\n#         )\n#         feature_info = self.backbone.feature_info.get_dicts()[-1]\n#         self.feature_dim = feature_info[\"num_chs\"]\n\n#         # C: Concept Head (Tạo CAMs)\n#         self.concept_head = nn.Conv2d(self.feature_dim, num_classes, kernel_size=1)\n\n#         # --- PHẦN 2: PROTOTYPES (Giai đoạn 2) ---\n#         self.embedding_dim = 128\n#         # P: Projector (Chiếu feature về không gian contrastive)\n#         self.projector = nn.Sequential(\n#             nn.Linear(self.feature_dim, 512),\n#             nn.ReLU(),\n#             nn.Linear(512, self.embedding_dim)\n#         )\n        \n#         # Learnable Prototypes: [Num_Classes, Num_Prototypes_Per_Class, Emb_Dim]\n#         self.prototypes = nn.Parameter(torch.randn(num_classes, num_prototypes, self.embedding_dim))\n        \n#         # --- PHẦN 3: TASK HEAD (Giai đoạn 3) ---\n#         # H: Task Head (Dự đoán bệnh từ điểm tương đồng)\n#         self.task_head = nn.Linear(num_classes * num_prototypes, num_classes)\n#         self.num_classes = num_classes\n#         self.num_prototypes = num_prototypes\n\n#     def get_features_and_cam(self, x):\n#         \"\"\"Dùng cho Giai đoạn 1\"\"\"\n#         if x.size(1) == 1:\n#             x = x.repeat(1, 3, 1, 1)\n#         features = self.backbone(x)[0]\n#         attn_logits = self.concept_head(features)\n#         return features, attn_logits\n\n#     def get_projected_vectors(self, features, attn_logits):\n#         \"\"\"Dùng cho Giai đoạn 2: Lấy Local Concept Vectors\"\"\"\n#         B, C, H, W = features.shape\n#         K = attn_logits.shape[1]\n        \n#         # 1. Normalize CAM (Spatial Softmax)\n#         attn_weights = F.softmax(attn_logits.view(B, K, -1), dim=-1).view(B, K, H, W)\n        \n#         # 2. Weighted Sum để lấy vector đại diện cho từng concept\n#         features_flat = features.view(B, C, -1).permute(0, 2, 1)\n#         local_concept_vectors = torch.bmm(attn_weights.view(B, K, -1), features_flat)\n        \n#         # 3. Project sang không gian embedding\n#         projected_vectors = self.projector(local_concept_vectors)\n#         return F.normalize(projected_vectors, p=2, dim=-1)\n\n#     def forward(self, x):\n#         \"\"\"Luồng chạy Full (Dùng cho Giai đoạn 3 & Inference)\"\"\"\n#         # 1. Trích xuất đặc trưng & CAM\n#         features, attn_logits = self.get_features_and_cam(x)\n        \n#         # 2. Tính projected vectors\n#         projected_vectors = self.get_projected_vectors(features, attn_logits)\n        \n#         # 3. Tính Similarity Score\n#         prototypes_norm = F.normalize(self.prototypes, p=2, dim=-1)\n#         sim_scores = torch.einsum('bkc,kmc->bkm', projected_vectors, prototypes_norm)\n#         s_vector = sim_scores.reshape(x.size(0), -1)\n        \n#         # 4. Predict\n#         logits = self.task_head(s_vector)\n        \n#         return {\n#             \"logits\": logits,\n#             \"attn_maps\": attn_logits,\n#             \"projected_vectors\": projected_vectors,\n#             \"sim_scores\": sim_scores\n#         }\n\n# print(\"✅ CSRModel defined!\")\n\n# # %% [markdown]\n# # ## Cell 4: Hàm tiền xử lý và visualize\n\n# # %%\n# def preprocess_image(image_path, target_size=384):\n#     \"\"\"Đọc và tiền xử lý ảnh\"\"\"\n#     if not os.path.exists(image_path):\n#         raise FileNotFoundError(f\"File not found: {image_path}\")\n        \n#     image = cv2.imread(image_path, cv2.IMREAD_GRAYSCALE)\n#     if image is None:\n#         raise ValueError(\"Cannot read image\")\n        \n#     image = cv2.resize(image, (target_size, target_size))\n#     img_norm = image.astype(np.float32) / 255.0\n#     img_tensor = torch.from_numpy(img_norm).unsqueeze(0).unsqueeze(0)\n#     return img_tensor, image\n\n\n# def visualize_result(original_img, probs, similarities, attn_maps, \n#                      class_names, top_k=3, save_path=None,\n#                      cam_threshold=0.3, show_binary_mask=False):\n#     \"\"\"\n#     Vẽ kết quả dự đoán với heatmap\n    \n#     Args:\n#         original_img: Ảnh gốc grayscale\n#         probs: Xác suất dự đoán [num_classes]\n#         similarities: Similarity scores [B, K, M]\n#         attn_maps: Attention maps [B, K, H, W]\n#         class_names: Danh sách tên các bệnh\n#         top_k: Số bệnh hiển thị\n#         save_path: Đường dẫn lưu ảnh\n#         cam_threshold: Ngưỡng để lọc CAM (0.0 - 1.0)\n#         show_binary_mask: True = mask nhị phân, False = gradient heatmap\n#     \"\"\"\n#     top_indices = np.argsort(probs)[::-1][:top_k]\n    \n#     fig, axes = plt.subplots(1, top_k + 1, figsize=(5 * (top_k + 1), 5))\n    \n#     # 1. Ảnh gốc\n#     axes[0].imshow(original_img, cmap='gray')\n#     axes[0].set_title(\"Input X-Ray\", fontsize=14)\n#     axes[0].axis('off')\n    \n#     info_text = f\"PREDICTIONS (CAM threshold={cam_threshold}):\\n\"\n#     for idx in top_indices:\n#         name = class_names[idx]\n#         sim_score = similarities[0, idx, :].max().item()\n#         prob = probs[idx]\n#         info_text += f\"{name}: {prob*100:.1f}% (Sim: {sim_score:.2f})\\n\"\n#     axes[0].set_xlabel(info_text, fontsize=10, loc='left')\n    \n#     # 2. Heatmap top-k với threshold\n#     for i, idx in enumerate(top_indices):\n#         name = class_names[idx]\n        \n#         # Lấy CAM\n#         cam = attn_maps[0, idx].cpu().numpy()\n#         cam_resized = cv2.resize(cam, (original_img.shape[1], original_img.shape[0]))\n        \n#         # Normalize về [0, 1]\n#         cam_norm = (cam_resized - cam_resized.min()) / (cam_resized.max() - cam_resized.min() + 1e-8)\n        \n#         # === ÁP DỤNG THRESHOLD ===\n#         if show_binary_mask:\n#             # Binary mask: chỉ hiện vùng > threshold\n#             cam_thresholded = (cam_norm > cam_threshold).astype(np.float32)\n#         else:\n#             # Gradient mask: giữ gradient nhưng zero-out vùng < threshold\n#             cam_thresholded = np.where(cam_norm > cam_threshold, cam_norm, 0)\n#             # Re-normalize sau khi threshold\n#             if cam_thresholded.max() > 0:\n#                 cam_thresholded = cam_thresholded / cam_thresholded.max()\n        \n#         # Tạo heatmap\n#         heatmap = cv2.applyColorMap(np.uint8(255 * cam_thresholded), cv2.COLORMAP_JET)\n#         heatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB)\n        \n#         # Overlay lên ảnh gốc\n#         img_rgb = cv2.cvtColor(original_img, cv2.COLOR_GRAY2RGB)\n        \n#         # Chỉ overlay vùng có giá trị (tránh màu xanh đậm ở vùng = 0)\n#         alpha_mask = (cam_thresholded > 0).astype(np.float32)\n#         alpha_mask = np.stack([alpha_mask] * 3, axis=-1)\n        \n#         overlay = img_rgb.copy().astype(np.float32)\n#         overlay = overlay * (1 - alpha_mask * 0.4) + heatmap.astype(np.float32) * (alpha_mask * 0.4)\n#         overlay = np.clip(overlay, 0, 255).astype(np.uint8)\n        \n#         axes[i + 1].imshow(overlay)\n#         axes[i + 1].set_title(f\"{name}\\n{probs[idx]*100:.1f}%\", fontsize=12)\n#         axes[i + 1].axis('off')\n    \n#     plt.tight_layout()\n    \n#     if save_path:\n#         plt.savefig(save_path, bbox_inches='tight', dpi=150)\n#         print(f\"💾 Result saved to: {save_path}\")\n    \n#     plt.show()\n\n\n# def visualize_with_multiple_thresholds(original_img, probs, similarities, attn_maps,\n#                                         class_names, class_idx=None, \n#                                         thresholds=[0.0, 0.3, 0.5, 0.7]):\n#     \"\"\"\n#     So sánh heatmap với nhiều mức threshold khác nhau\n    \n#     Args:\n#         class_idx: Index của class cần visualize (None = top-1)\n#         thresholds: List các giá trị threshold để so sánh\n#     \"\"\"\n#     if class_idx is None:\n#         class_idx = np.argmax(probs)\n    \n#     name = class_names[class_idx]\n#     prob = probs[class_idx]\n    \n#     fig, axes = plt.subplots(1, len(thresholds) + 1, figsize=(4 * (len(thresholds) + 1), 4))\n    \n#     # Ảnh gốc\n#     axes[0].imshow(original_img, cmap='gray')\n#     axes[0].set_title(f\"Original\\n{name}: {prob*100:.1f}%\", fontsize=11)\n#     axes[0].axis('off')\n    \n#     # CAM với các threshold khác nhau\n#     cam = attn_maps[0, class_idx].cpu().numpy()\n#     cam_resized = cv2.resize(cam, (original_img.shape[1], original_img.shape[0]))\n#     cam_norm = (cam_resized - cam_resized.min()) / (cam_resized.max() - cam_resized.min() + 1e-8)\n    \n#     for i, thresh in enumerate(thresholds):\n#         cam_thresholded = np.where(cam_norm > thresh, cam_norm, 0)\n#         if cam_thresholded.max() > 0:\n#             cam_thresholded = cam_thresholded / cam_thresholded.max()\n        \n#         heatmap = cv2.applyColorMap(np.uint8(255 * cam_thresholded), cv2.COLORMAP_JET)\n#         heatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB)\n        \n#         img_rgb = cv2.cvtColor(original_img, cv2.COLOR_GRAY2RGB)\n        \n#         alpha_mask = (cam_thresholded > 0).astype(np.float32)\n#         alpha_mask = np.stack([alpha_mask] * 3, axis=-1)\n        \n#         overlay = img_rgb.copy().astype(np.float32)\n#         overlay = overlay * (1 - alpha_mask * 0.5) + heatmap.astype(np.float32) * (alpha_mask * 0.5)\n#         overlay = np.clip(overlay, 0, 255).astype(np.uint8)\n        \n#         # Tính % vùng được highlight\n#         coverage = (cam_norm > thresh).sum() / cam_norm.size * 100\n        \n#         axes[i + 1].imshow(overlay)\n#         axes[i + 1].set_title(f\"Threshold={thresh}\\nCoverage: {coverage:.1f}%\", fontsize=11)\n#         axes[i + 1].axis('off')\n    \n#     plt.suptitle(f\"CAM Threshold Comparison: {name}\", fontsize=14)\n#     plt.tight_layout()\n#     plt.show()\n\n\n# def visualize_all_classes(original_img, probs, attn_maps, class_names, \n#                           cam_threshold=0.3, prob_threshold=0.1):\n#     \"\"\"\n#     Hiển thị heatmap cho TẤT CẢ các class có probability > prob_threshold\n    \n#     Args:\n#         prob_threshold: Chỉ hiển thị class có prob > threshold này\n#     \"\"\"\n#     # Lọc các class có prob > threshold\n#     valid_indices = np.where(probs > prob_threshold)[0]\n    \n#     if len(valid_indices) == 0:\n#         print(f\"⚠️ Không có class nào có probability > {prob_threshold}\")\n#         return\n    \n#     # Sắp xếp theo prob giảm dần\n#     valid_indices = valid_indices[np.argsort(probs[valid_indices])[::-1]]\n    \n#     n_classes = len(valid_indices)\n#     n_cols = min(4, n_classes + 1)\n#     n_rows = (n_classes + n_cols) // n_cols\n    \n#     fig, axes = plt.subplots(n_rows, n_cols, figsize=(4 * n_cols, 4 * n_rows))\n#     axes = axes.flatten() if n_rows > 1 else [axes] if n_cols == 1 else axes.flatten()\n    \n#     # Ảnh gốc\n#     axes[0].imshow(original_img, cmap='gray')\n#     axes[0].set_title(\"Original X-Ray\", fontsize=12)\n#     axes[0].axis('off')\n    \n#     # Heatmap cho từng class\n#     for i, idx in enumerate(valid_indices):\n#         if i + 1 >= len(axes):\n#             break\n            \n#         name = class_names[idx]\n#         prob = probs[idx]\n        \n#         cam = attn_maps[0, idx].cpu().numpy()\n#         cam_resized = cv2.resize(cam, (original_img.shape[1], original_img.shape[0]))\n#         cam_norm = (cam_resized - cam_resized.min()) / (cam_resized.max() - cam_resized.min() + 1e-8)\n        \n#         cam_thresholded = np.where(cam_norm > cam_threshold, cam_norm, 0)\n#         if cam_thresholded.max() > 0:\n#             cam_thresholded = cam_thresholded / cam_thresholded.max()\n        \n#         heatmap = cv2.applyColorMap(np.uint8(255 * cam_thresholded), cv2.COLORMAP_JET)\n#         heatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB)\n        \n#         img_rgb = cv2.cvtColor(original_img, cv2.COLOR_GRAY2RGB)\n        \n#         alpha_mask = (cam_thresholded > 0).astype(np.float32)\n#         alpha_mask = np.stack([alpha_mask] * 3, axis=-1)\n        \n#         overlay = img_rgb.copy().astype(np.float32)\n#         overlay = overlay * (1 - alpha_mask * 0.4) + heatmap.astype(np.float32) * (alpha_mask * 0.4)\n#         overlay = np.clip(overlay, 0, 255).astype(np.uint8)\n        \n#         axes[i + 1].imshow(overlay)\n#         axes[i + 1].set_title(f\"{name}\\n{prob*100:.1f}%\", fontsize=10)\n#         axes[i + 1].axis('off')\n    \n#     # Ẩn các axes thừa\n#     for j in range(i + 2, len(axes)):\n#         axes[j].axis('off')\n    \n#     plt.suptitle(f\"All Detected Conditions (prob > {prob_threshold*100:.0f}%)\", fontsize=14)\n#     plt.tight_layout()\n#     plt.show()\n\n# print(\"✅ Helper functions defined!\")\n\n# # %% [markdown]\n# # ## Cell 5: Load Model\n\n# # %%\n# device = torch.device(DEVICE if torch.cuda.is_available() else \"cpu\")\n# print(f\"🖥️  Device: {device}\")\n\n# # Tạo model\n# model = CSRModel(\n#     num_classes=NUM_CLASSES,\n#     num_prototypes=NUM_PROTOTYPES,\n#     model_name=MODEL_NAME,\n#     pretrained=False  # Không cần pretrained vì sẽ load checkpoint\n# )\n\n# # Load checkpoint\n# print(f\"📂 Loading checkpoint: {CHECKPOINT_PATH}\")\n# ckpt = torch.load(CHECKPOINT_PATH, map_location=device)\n\n# if 'model_state_dict' in ckpt:\n#     model.load_state_dict(ckpt['model_state_dict'])\n#     print(f\"   Epoch: {ckpt.get('epoch', 'N/A')}\")\n# elif 'state_dict' in ckpt:\n#     model.load_state_dict(ckpt['state_dict'])\n# else:\n#     model.load_state_dict(ckpt)\n\n# model.to(device)\n# model.eval()\n# print(\"✅ Model loaded successfully!\")\n\n# # %% [markdown]\n# # ## Cell 6: Inference trên 1 ảnh\n\n# # %%\n# # Tiền xử lý\n# print(f\"🖼️  Loading image: {IMAGE_PATH}\")\n# img_tensor, original_img = preprocess_image(IMAGE_PATH, target_size=IMG_SIZE)\n# img_tensor = img_tensor.to(device)\n\n# # Inference\n# print(\"🔮 Running inference...\")\n# with torch.no_grad():\n#     outputs = model(img_tensor)\n    \n#     logits = outputs['logits'][0]\n#     sim_scores = outputs['sim_scores']\n#     attn_maps = outputs['attn_maps']\n    \n#     probs = torch.sigmoid(logits).cpu().numpy()\n\n# # Hiển thị kết quả dạng text\n# print(\"\\n📊 Prediction Results:\")\n# print(\"-\" * 50)\n# for i, (name, prob) in enumerate(zip(CLASS_NAMES, probs)):\n#     bar = \"█\" * int(prob * 20)\n#     status = \"⚠️\" if prob > 0.5 else \"  \"\n#     print(f\"{status} {name:20s}: {prob*100:5.1f}% {bar}\")\n\n# # Hiển thị các bệnh được phát hiện (prob > 0.5)\n# detected = [(CLASS_NAMES[i], probs[i]) for i in range(len(probs)) if probs[i] > 0.5]\n# print(\"\\n\" + \"=\" * 50)\n# if detected:\n#     print(\"🔴 DETECTED CONDITIONS:\")\n#     for name, prob in detected:\n#         print(f\"   • {name}: {prob*100:.1f}%\")\n# else:\n#     print(\"🟢 No significant conditions detected (all < 50%)\")\n\n# # %% [markdown]\n# # ## Cell 7: Visualize Top-K với Threshold\n\n# # %%\n# # Visualize với threshold\n# visualize_result(\n#     original_img, probs, sim_scores, attn_maps,\n#     class_names=CLASS_NAMES,\n#     top_k=TOP_K,\n#     save_path=SAVE_PATH,\n#     cam_threshold=CAM_THRESHOLD,\n#     show_binary_mask=SHOW_BINARY_MASK\n# )\n\n# # %% [markdown]\n# # ## Cell 8: So sánh các mức Threshold\n\n# # %%\n# # So sánh nhiều threshold cho bệnh có xác suất cao nhất\n# visualize_with_multiple_thresholds(\n#     original_img, probs, sim_scores, attn_maps,\n#     class_names=CLASS_NAMES,\n#     class_idx=None,  # None = tự động chọn top-1, hoặc đặt số 0-13\n#     thresholds=[0.0, 0.2, 0.4, 0.6, 0.8]\n# )\n\n# # %% [markdown]\n# # ## Cell 9: Hiển thị tất cả các bệnh phát hiện được\n\n# # %%\n# # Hiển thị heatmap cho tất cả class có probability > 10%\n# visualize_all_classes(\n#     original_img, probs, attn_maps, \n#     class_names=CLASS_NAMES,\n#     cam_threshold=CAM_THRESHOLD,\n#     prob_threshold=0.6  # Chỉ hiện class có prob > 10%\n# )\n\n# # %% [markdown]\n# # ## Cell 10 (Optional): Inference trên nhiều ảnh\n\n# # %%\n# def batch_inference(image_paths, model, device, class_names, img_size=384, prob_threshold=0.5):\n#     \"\"\"Inference trên nhiều ảnh và trả về kết quả\"\"\"\n#     results = []\n    \n#     for path in image_paths:\n#         try:\n#             img_tensor, original_img = preprocess_image(path, target_size=img_size)\n#             img_tensor = img_tensor.to(device)\n            \n#             with torch.no_grad():\n#                 outputs = model(img_tensor)\n#                 probs = torch.sigmoid(outputs['logits'][0]).cpu().numpy()\n            \n#             # Lấy các bệnh có prob > threshold\n#             detected = [(class_names[i], probs[i]) for i in range(len(probs)) if probs[i] > prob_threshold]\n            \n#             results.append({\n#                 'path': path,\n#                 'probs': probs,\n#                 'detected': detected,\n#                 'top_pred': class_names[np.argmax(probs)],\n#                 'top_prob': probs.max()\n#             })\n#         except Exception as e:\n#             results.append({\n#                 'path': path,\n#                 'error': str(e)\n#             })\n    \n#     return results\n\n# # Ví dụ sử dụng:\n# # IMAGE_LIST = [\n# #     \"/kaggle/input/images/img1.png\",\n# #     \"/kaggle/input/images/img2.png\",\n# #     \"/kaggle/input/images/img3.png\",\n# # ]\n# # results = batch_inference(IMAGE_LIST, model, device, CLASS_NAMES, IMG_SIZE)\n# # \n# # for r in results:\n# #     if 'error' in r:\n# #         print(f\"❌ {r['path']}: {r['error']}\")\n# #     else:\n# #         print(f\"✅ {r['path']}:\")\n# #         if r['detected']:\n# #             for name, prob in r['detected']:\n# #                 print(f\"   • {name}: {prob*100:.1f}%\")\n# #         else:\n# #             print(f\"   No conditions detected\")\n\n# # %% [markdown]\n# # ## Cell 11 (Optional): Interactive Threshold Slider (Jupyter)\n\n# # %%\n# # Chỉ chạy được trong Jupyter Notebook với ipywidgets\n# try:\n#     from ipywidgets import interact, FloatSlider\n    \n#     def interactive_threshold(threshold=0.3):\n#         visualize_result(\n#             original_img, probs, sim_scores, attn_maps,\n#             class_names=CLASS_NAMES,\n#             top_k=TOP_K,\n#             save_path=None,\n#             cam_threshold=threshold,\n#             show_binary_mask=False\n#         )\n    \n#     interact(\n#         interactive_threshold,\n#         threshold=FloatSlider(min=0.0, max=1.0, step=0.05, value=0.3, description='CAM Threshold')\n#     )\n# except ImportError:\n#     print(\"⚠️ ipywidgets not available. Skip interactive slider.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T15:39:30.174559Z","iopub.execute_input":"2025-12-10T15:39:30.174955Z","iopub.status.idle":"2025-12-10T15:39:33.659394Z","shell.execute_reply.started":"2025-12-10T15:39:30.174919Z","shell.execute_reply":"2025-12-10T15:39:33.658827Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# \"\"\"\n# Test model SAU Phase 1 - Chỉ dùng concept_head, KHÔNG dùng task_head\n# \"\"\"\n\n# import torch\n# import torch.nn as nn\n# import torch.nn.functional as F\n# import numpy as np\n# import cv2\n# import os\n# import timm\n# import pandas as pd\n\n# # =================== CẤU HÌNH ===================\n# CHECKPOINT_PATH = \"/kaggle/input/vin-csr-training/checkpoints/csr_phase1.pth\"\n# IMAGE_PATH = \"/kaggle/input/vindr-image-convert/train_png_384/9a5094b2563a1ef3ff50dc5c7ff71345.png\"\n# CSV_PATH = \"/kaggle/input/png-trainer-9/pytorch/default/1/src/csv/test_split.csv\"\n# IMAGE_DIR = \"/kaggle/input/vindr-image-convert/train_png_384\"\n\n# NUM_CLASSES = 14\n# NUM_PROTOTYPES = 15\n# MODEL_NAME = \"densenet121\"\n# IMG_SIZE = 384\n# DEVICE = \"cuda\"\n\n# CLASS_NAMES = [\n#     'Aortic enlargement', 'Atelectasis', 'Calcification', 'Cardiomegaly',\n#     'Consolidation', 'ILD', 'Infiltration', 'Lung Opacity',\n#     'Nodule/Mass', 'Other lesion', 'Pleural effusion', 'Pleural thickening',\n#     'Pneumothorax', 'Pulmonary fibrosis'\n# ]\n\n# # ================================================\n\n# class CSRModel(nn.Module):\n#     def __init__(self, num_classes=14, num_prototypes=5, model_name=\"resnet50\", pretrained=True):\n#         super().__init__()\n        \n#         self.backbone = timm.create_model(\n#             model_name, pretrained=pretrained, features_only=True, out_indices=(4,)\n#         )\n#         feature_info = self.backbone.feature_info.get_dicts()[-1]\n#         self.feature_dim = feature_info[\"num_chs\"]\n#         self.concept_head = nn.Conv2d(self.feature_dim, num_classes, kernel_size=1)\n        \n#         self.embedding_dim = 128\n#         self.projector = nn.Sequential(\n#             nn.Linear(self.feature_dim, 512),\n#             nn.ReLU(),\n#             nn.Linear(512, self.embedding_dim)\n#         )\n        \n#         self.prototypes = nn.Parameter(torch.randn(num_classes, num_prototypes, self.embedding_dim))\n#         self.task_head = nn.Linear(num_classes * num_prototypes, num_classes)\n#         self.num_classes = num_classes\n#         self.num_prototypes = num_prototypes\n\n#     def get_features_and_cam(self, x):\n#         if x.size(1) == 1:\n#             x = x.repeat(1, 3, 1, 1)\n#         features = self.backbone(x)[0]\n#         attn_logits = self.concept_head(features)\n#         return features, attn_logits\n    \n#     def forward_phase1(self, x):\n#         \"\"\"\n#         Forward CHỈ DÙNG CHO PHASE 1\n#         Dự đoán bằng Global Average Pooling trên CAM\n#         \"\"\"\n#         _, attn_logits = self.get_features_and_cam(x)\n#         # GAP: [B, K, H, W] -> [B, K]\n#         logits = F.adaptive_avg_pool2d(attn_logits, (1, 1)).view(x.size(0), -1)\n#         return {\n#             \"logits\": logits,\n#             \"attn_maps\": attn_logits\n#         }\n\n#     def forward(self, x):\n#         # Full forward (Phase 3)\n#         features, attn_logits = self.get_features_and_cam(x)\n#         projected_vectors = self.get_projected_vectors(features, attn_logits)\n#         prototypes_norm = F.normalize(self.prototypes, p=2, dim=-1)\n#         sim_scores = torch.einsum('bkc,kmc->bkm', projected_vectors, prototypes_norm)\n#         s_vector = sim_scores.reshape(x.size(0), -1)\n#         logits = self.task_head(s_vector)\n#         return {\"logits\": logits, \"attn_maps\": attn_logits, \"sim_scores\": sim_scores}\n    \n#     def get_projected_vectors(self, features, attn_logits):\n#         B, C, H, W = features.shape\n#         K = attn_logits.shape[1]\n#         attn_weights = F.softmax(attn_logits.view(B, K, -1), dim=-1).view(B, K, H, W)\n#         features_flat = features.view(B, C, -1).permute(0, 2, 1)\n#         local_concept_vectors = torch.bmm(attn_weights.view(B, K, -1), features_flat)\n#         projected_vectors = self.projector(local_concept_vectors)\n#         return F.normalize(projected_vectors, p=2, dim=-1)\n\n\n# def preprocess_image(image_path, target_size=384):\n#     image = cv2.imread(image_path, cv2.IMREAD_GRAYSCALE)\n#     if image is None:\n#         raise ValueError(f\"Cannot read image: {image_path}\")\n#     image = cv2.resize(image, (target_size, target_size))\n#     img_norm = image.astype(np.float32) / 255.0\n#     img_tensor = torch.from_numpy(img_norm).unsqueeze(0).unsqueeze(0)\n#     return img_tensor\n\n\n# def predict_phase1(model, image_path, device, threshold=0.5):\n#     \"\"\"\n#     Dự đoán ĐÚNG CÁCH cho Phase 1\n#     \"\"\"\n#     img_tensor = preprocess_image(image_path, IMG_SIZE)\n#     img_tensor = img_tensor.to(device)\n    \n#     with torch.no_grad():\n#         # SỬ DỤNG forward_phase1 thay vì forward\n#         outputs = model.forward_phase1(img_tensor)\n#         logits = outputs['logits'][0]\n#         probs = torch.sigmoid(logits).cpu().numpy()\n    \n#     predictions = (probs > threshold).astype(int)\n#     return probs, predictions\n\n\n# def get_ground_truth(image_id, csv_path):\n#     df = pd.read_csv(csv_path)\n#     rows = df[df['image_id'] == image_id]\n    \n#     gt_labels = np.zeros(NUM_CLASSES)\n#     for _, row in rows.iterrows():\n#         class_id = row['class_id']\n#         if 0 <= class_id < NUM_CLASSES:\n#             gt_labels[class_id] = 1\n#     return gt_labels\n\n\n# def print_results(image_id, probs, predictions, gt_labels):\n#     print(\"\\n\" + \"=\" * 70)\n#     print(f\"📷 Image: {image_id}\")\n#     print(\"=\" * 70)\n#     print(f\"{'Class':<25} {'Prob':>8} {'Pred':>6} {'GT':>6} {'Match':>7}\")\n#     print(\"-\" * 70)\n    \n#     correct = 0\n#     for i in range(NUM_CLASSES):\n#         prob = probs[i]\n#         pred = predictions[i]\n#         gt = int(gt_labels[i])\n#         match = \"✅\" if pred == gt else \"❌\"\n#         if pred == gt:\n#             correct += 1\n#         prefix = \"🔴\" if pred == 1 else \"  \"\n#         print(f\"{prefix}{CLASS_NAMES[i]:<23} {prob*100:>7.1f}% {pred:>6} {gt:>6} {match:>7}\")\n    \n#     print(\"-\" * 70)\n#     pred_diseases = [CLASS_NAMES[i] for i in range(NUM_CLASSES) if predictions[i] == 1]\n#     gt_diseases = [CLASS_NAMES[i] for i in range(NUM_CLASSES) if gt_labels[i] == 1]\n    \n#     print(f\"\\n🔴 Predicted: {pred_diseases if pred_diseases else 'None'}\")\n#     print(f\"🟢 Ground truth: {gt_diseases if gt_diseases else 'None'}\")\n#     print(f\"\\n📊 Accuracy: {correct}/{NUM_CLASSES} ({correct/NUM_CLASSES*100:.1f}%)\")\n    \n#     if np.array_equal(predictions, gt_labels.astype(int)):\n#         print(\"✅ EXACT MATCH!\")\n#     else:\n#         print(\"❌ NOT EXACT MATCH\")\n\n\n# def test_multiple_samples(model, csv_path, image_dir, device, num_samples=10):\n#     \"\"\"Test trên nhiều mẫu\"\"\"\n#     df = pd.read_csv(csv_path)\n#     unique_images = df['image_id'].unique()\n#     selected = np.random.choice(unique_images, min(num_samples, len(unique_images)), replace=False)\n    \n#     print(f\"\\n🧪 Testing {len(selected)} samples with PHASE 1 forward...\")\n    \n#     results = []\n#     for image_id in selected:\n#         image_path = os.path.join(image_dir, f\"{image_id}.png\")\n#         if not os.path.exists(image_path):\n#             continue\n        \n#         probs, predictions = predict_phase1(model, image_path, device)\n#         gt_labels = get_ground_truth(image_id, csv_path)\n#         print_results(image_id, probs, predictions, gt_labels)\n        \n#         exact_match = np.array_equal(predictions, gt_labels.astype(int))\n#         element_acc = (predictions == gt_labels.astype(int)).mean()\n#         results.append({'exact_match': exact_match, 'element_acc': element_acc})\n    \n#     if results:\n#         print(\"\\n\" + \"=\" * 70)\n#         print(\"📈 SUMMARY\")\n#         print(\"=\" * 70)\n#         exact_matches = sum(r['exact_match'] for r in results)\n#         avg_acc = np.mean([r['element_acc'] for r in results])\n#         print(f\"Exact Match: {exact_matches}/{len(results)} ({exact_matches/len(results)*100:.1f}%)\")\n#         print(f\"Avg Element Accuracy: {avg_acc*100:.1f}%\")\n\n\n# # =================== MAIN ===================\n# if __name__ == \"__main__\":\n#     device = torch.device(DEVICE if torch.cuda.is_available() else \"cpu\")\n#     print(f\"🖥️ Device: {device}\")\n    \n#     # Load model\n#     print(f\"\\n📂 Loading model from: {CHECKPOINT_PATH}\")\n#     model = CSRModel(\n#         num_classes=NUM_CLASSES,\n#         num_prototypes=NUM_PROTOTYPES,\n#         model_name=MODEL_NAME,\n#         pretrained=False\n#     )\n    \n#     ckpt = torch.load(CHECKPOINT_PATH, map_location=device)\n#     if 'model_state_dict' in ckpt:\n#         model.load_state_dict(ckpt['model_state_dict'])\n#     else:\n#         model.load_state_dict(ckpt)\n    \n#     model.to(device)\n#     model.eval()\n#     print(\"✅ Model loaded!\")\n    \n#     # Test single image\n#     print(f\"\\n🔍 Testing: {IMAGE_PATH}\")\n#     probs, predictions = predict_phase1(model, IMAGE_PATH, device)\n    \n#     image_id = os.path.splitext(os.path.basename(IMAGE_PATH))[0]\n#     gt_labels = get_ground_truth(image_id, CSV_PATH)\n#     print_results(image_id, probs, predictions, gt_labels)\n    \n#     # Test multiple samples\n#     # test_multiple_samples(model, CSV_PATH, IMAGE_DIR, device, num_samples=10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T15:49:25.219902Z","iopub.execute_input":"2025-12-10T15:49:25.220489Z","iopub.status.idle":"2025-12-10T15:49:25.89192Z","shell.execute_reply.started":"2025-12-10T15:49:25.220465Z","shell.execute_reply":"2025-12-10T15:49:25.891073Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# \"\"\"\n# Đánh giá mô hình CSR Phase 1 trên toàn bộ tập test\n# Metrics: Accuracy, Precision, Recall, F1, AUC\n# \"\"\"\n\n# import torch\n# import torch.nn as nn\n# import torch.nn.functional as F\n# import numpy as np\n# import cv2\n# import os\n# import timm\n# import pandas as pd\n# from tqdm import tqdm\n# from sklearn.metrics import (\n#     accuracy_score, precision_score, recall_score, f1_score,\n#     roc_auc_score, classification_report, multilabel_confusion_matrix\n# )\n\n# # =================== CẤU HÌNH ===================\n# CHECKPOINT_PATH = \"/kaggle/input/vin-csr-training/checkpoints/csr_phase1.pth\"\n# CSV_PATH = \"/kaggle/input/png-trainer-9/pytorch/default/1/src/csv/test_split.csv\"\n# IMAGE_DIR = \"/kaggle/input/vindr-image-convert/train_png_384\"\n\n# NUM_CLASSES = 14\n# NUM_PROTOTYPES = 15\n# MODEL_NAME = \"densenet121\"\n# IMG_SIZE = 384\n# DEVICE = \"cuda\"\n# THRESHOLD = 0.5\n# BATCH_SIZE = 32\n\n# CLASS_NAMES = [\n#     'Aortic enlargement', 'Atelectasis', 'Calcification', 'Cardiomegaly',\n#     'Consolidation', 'ILD', 'Infiltration', 'Lung Opacity',\n#     'Nodule/Mass', 'Other lesion', 'Pleural effusion', 'Pleural thickening',\n#     'Pneumothorax', 'Pulmonary fibrosis'\n# ]\n\n# # =================== MODEL ===================\n# class CSRModel(nn.Module):\n#     def __init__(self, num_classes=14, num_prototypes=5, model_name=\"resnet50\", pretrained=True):\n#         super().__init__()\n#         self.backbone = timm.create_model(\n#             model_name, pretrained=pretrained, features_only=True, out_indices=(4,)\n#         )\n#         feature_info = self.backbone.feature_info.get_dicts()[-1]\n#         self.feature_dim = feature_info[\"num_chs\"]\n#         self.concept_head = nn.Conv2d(self.feature_dim, num_classes, kernel_size=1)\n        \n#         self.embedding_dim = 128\n#         self.projector = nn.Sequential(\n#             nn.Linear(self.feature_dim, 512),\n#             nn.ReLU(),\n#             nn.Linear(512, self.embedding_dim)\n#         )\n#         self.prototypes = nn.Parameter(torch.randn(num_classes, num_prototypes, self.embedding_dim))\n#         self.task_head = nn.Linear(num_classes * num_prototypes, num_classes)\n#         self.num_classes = num_classes\n#         self.num_prototypes = num_prototypes\n\n#     def forward_phase1(self, x):\n#         if x.size(1) == 1:\n#             x = x.repeat(1, 3, 1, 1)\n#         features = self.backbone(x)[0]\n#         attn_logits = self.concept_head(features)\n#         logits = F.adaptive_avg_pool2d(attn_logits, (1, 1)).view(x.size(0), -1)\n#         return logits\n\n\n# # =================== DATA LOADING ===================\n# def load_test_data(csv_path, image_dir):\n#     \"\"\"Load toàn bộ test data\"\"\"\n#     df = pd.read_csv(csv_path)\n    \n#     # Tạo ground truth matrix\n#     unique_images = df['image_id'].unique()\n#     print(f\"📊 Tổng số ảnh test: {len(unique_images)}\")\n    \n#     image_labels = {}\n#     for image_id in unique_images:\n#         labels = np.zeros(NUM_CLASSES)\n#         rows = df[df['image_id'] == image_id]\n#         for _, row in rows.iterrows():\n#             class_id = row['class_id']\n#             if 0 <= class_id < NUM_CLASSES:\n#                 labels[class_id] = 1\n#         image_labels[image_id] = labels\n    \n#     # Lọc ảnh tồn tại\n#     valid_images = []\n#     for image_id in unique_images:\n#         image_path = os.path.join(image_dir, f\"{image_id}.png\")\n#         if os.path.exists(image_path):\n#             valid_images.append(image_id)\n    \n#     print(f\"✅ Số ảnh hợp lệ: {len(valid_images)}\")\n#     return valid_images, image_labels\n\n\n# def preprocess_image(image_path, target_size=384):\n#     \"\"\"Preprocess single image\"\"\"\n#     image = cv2.imread(image_path, cv2.IMREAD_GRAYSCALE)\n#     if image is None:\n#         return None\n#     image = cv2.resize(image, (target_size, target_size))\n#     img_norm = image.astype(np.float32) / 255.0\n#     img_tensor = torch.from_numpy(img_norm).unsqueeze(0).unsqueeze(0)\n#     return img_tensor\n\n\n# # =================== EVALUATION ===================\n# def evaluate_model(model, image_ids, image_labels, image_dir, device, threshold=0.5):\n#     \"\"\"\n#     Đánh giá model trên toàn bộ tập test\n#     \"\"\"\n#     model.eval()\n    \n#     all_probs = []\n#     all_preds = []\n#     all_labels = []\n    \n#     print(\"\\n🔄 Đang dự đoán...\")\n#     with torch.no_grad():\n#         for image_id in tqdm(image_ids, desc=\"Evaluating\"):\n#             image_path = os.path.join(image_dir, f\"{image_id}.png\")\n#             img_tensor = preprocess_image(image_path, IMG_SIZE)\n            \n#             if img_tensor is None:\n#                 continue\n            \n#             img_tensor = img_tensor.to(device)\n#             logits = model.forward_phase1(img_tensor)\n#             probs = torch.sigmoid(logits).cpu().numpy()[0]\n#             preds = (probs > threshold).astype(int)\n            \n#             all_probs.append(probs)\n#             all_preds.append(preds)\n#             all_labels.append(image_labels[image_id])\n    \n#     all_probs = np.array(all_probs)\n#     all_preds = np.array(all_preds)\n#     all_labels = np.array(all_labels)\n    \n#     return all_probs, all_preds, all_labels\n\n\n# def compute_metrics(all_probs, all_preds, all_labels):\n#     \"\"\"\n#     Tính toán các metrics đánh giá\n#     \"\"\"\n#     results = {}\n    \n#     # ===== 1. SAMPLE-WISE METRICS =====\n#     # Exact Match Ratio (Subset Accuracy)\n#     exact_match = np.all(all_preds == all_labels, axis=1).mean()\n#     results['exact_match_ratio'] = exact_match\n    \n#     # Element-wise Accuracy\n#     element_acc = (all_preds == all_labels).mean()\n#     results['element_accuracy'] = element_acc\n    \n#     # ===== 2. LABEL-WISE METRICS (Macro/Micro) =====\n#     # Precision, Recall, F1\n#     results['precision_macro'] = precision_score(all_labels, all_preds, average='macro', zero_division=0)\n#     results['precision_micro'] = precision_score(all_labels, all_preds, average='micro', zero_division=0)\n#     results['recall_macro'] = recall_score(all_labels, all_preds, average='macro', zero_division=0)\n#     results['recall_micro'] = recall_score(all_labels, all_preds, average='micro', zero_division=0)\n#     results['f1_macro'] = f1_score(all_labels, all_preds, average='macro', zero_division=0)\n#     results['f1_micro'] = f1_score(all_labels, all_preds, average='micro', zero_division=0)\n    \n#     # ===== 3. AUC-ROC =====\n#     try:\n#         results['auc_macro'] = roc_auc_score(all_labels, all_probs, average='macro')\n#         results['auc_micro'] = roc_auc_score(all_labels, all_probs, average='micro')\n#     except ValueError as e:\n#         print(f\"⚠️ Cannot compute AUC: {e}\")\n#         results['auc_macro'] = None\n#         results['auc_micro'] = None\n    \n#     # ===== 4. PER-CLASS METRICS =====\n#     per_class = {}\n#     for i, class_name in enumerate(CLASS_NAMES):\n#         y_true = all_labels[:, i]\n#         y_pred = all_preds[:, i]\n#         y_prob = all_probs[:, i]\n        \n#         # Đếm số lượng\n#         n_positive = int(y_true.sum())\n#         n_pred_positive = int(y_pred.sum())\n        \n#         # Tính metrics\n#         acc = accuracy_score(y_true, y_pred)\n#         prec = precision_score(y_true, y_pred, zero_division=0)\n#         rec = recall_score(y_true, y_pred, zero_division=0)\n#         f1 = f1_score(y_true, y_pred, zero_division=0)\n        \n#         try:\n#             auc = roc_auc_score(y_true, y_prob) if n_positive > 0 and n_positive < len(y_true) else None\n#         except:\n#             auc = None\n        \n#         per_class[class_name] = {\n#             'n_positive': n_positive,\n#             'n_pred_positive': n_pred_positive,\n#             'accuracy': acc,\n#             'precision': prec,\n#             'recall': rec,\n#             'f1': f1,\n#             'auc': auc\n#         }\n    \n#     results['per_class'] = per_class\n#     return results\n\n\n# def print_evaluation_results(results, all_labels):\n#     \"\"\"In kết quả đánh giá\"\"\"\n    \n#     print(\"\\n\" + \"=\" * 80)\n#     print(\"📊 KẾT QUẢ ĐÁNH GIÁ MÔ HÌNH CSR PHASE 1\")\n#     print(\"=\" * 80)\n    \n#     print(f\"\\n📈 OVERALL METRICS:\")\n#     print(\"-\" * 50)\n#     print(f\"  Exact Match Ratio:     {results['exact_match_ratio']*100:>6.2f}%\")\n#     print(f\"  Element Accuracy:      {results['element_accuracy']*100:>6.2f}%\")\n#     print(f\"  Precision (Macro):     {results['precision_macro']*100:>6.2f}%\")\n#     print(f\"  Precision (Micro):     {results['precision_micro']*100:>6.2f}%\")\n#     print(f\"  Recall (Macro):        {results['recall_macro']*100:>6.2f}%\")\n#     print(f\"  Recall (Micro):        {results['recall_micro']*100:>6.2f}%\")\n#     print(f\"  F1-Score (Macro):      {results['f1_macro']*100:>6.2f}%\")\n#     print(f\"  F1-Score (Micro):      {results['f1_micro']*100:>6.2f}%\")\n#     if results['auc_macro']:\n#         print(f\"  AUC-ROC (Macro):       {results['auc_macro']*100:>6.2f}%\")\n#         print(f\"  AUC-ROC (Micro):       {results['auc_micro']*100:>6.2f}%\")\n    \n#     print(f\"\\n📋 PER-CLASS METRICS:\")\n#     print(\"-\" * 100)\n#     header = f\"{'Class':<25} {'N_GT':>6} {'N_Pred':>7} {'Acc':>7} {'Prec':>7} {'Rec':>7} {'F1':>7} {'AUC':>7}\"\n#     print(header)\n#     print(\"-\" * 100)\n    \n#     per_class = results['per_class']\n#     for class_name in CLASS_NAMES:\n#         m = per_class[class_name]\n#         auc_str = f\"{m['auc']*100:>6.1f}%\" if m['auc'] else \"   N/A\"\n#         print(f\"{class_name:<25} {m['n_positive']:>6} {m['n_pred_positive']:>7} \"\n#               f\"{m['accuracy']*100:>6.1f}% {m['precision']*100:>6.1f}% \"\n#               f\"{m['recall']*100:>6.1f}% {m['f1']*100:>6.1f}% {auc_str}\")\n    \n#     print(\"-\" * 100)\n    \n#     # Thống kê thêm\n#     print(f\"\\n📊 THỐNG KÊ BỔ SUNG:\")\n#     print(\"-\" * 50)\n#     total_positive = int(all_labels.sum())\n#     total_samples = len(all_labels)\n#     print(f\"  Tổng số mẫu test:      {total_samples}\")\n#     print(f\"  Tổng số nhãn dương:    {total_positive}\")\n#     print(f\"  Trung bình nhãn/ảnh:   {total_positive/total_samples:.2f}\")\n\n\n# def compute_simple_accuracy(all_preds, all_labels):\n#     \"\"\"\n#     Đơn giản: Đếm số nhãn dự đoán đúng\n#     \"\"\"\n#     # Tổng số nhãn đúng\n#     correct = (all_preds == all_labels).sum()\n#     total = all_preds.size\n    \n#     # Đếm theo từng loại\n#     true_positive = ((all_preds == 1) & (all_labels == 1)).sum()\n#     true_negative = ((all_preds == 0) & (all_labels == 0)).sum()\n#     false_positive = ((all_preds == 1) & (all_labels == 0)).sum()\n#     false_negative = ((all_preds == 0) & (all_labels == 1)).sum()\n    \n#     print(\"\\n\" + \"=\" * 60)\n#     print(\"🎯 ĐÁNH GIÁ ĐƠN GIẢN: SỐ LƯỢNG NHÃN DỰ ĐOÁN ĐÚNG\")\n#     print(\"=\" * 60)\n#     print(f\"  ✅ Tổng số nhãn đúng:           {correct:,} / {total:,}\")\n#     print(f\"  📊 Accuracy:                    {correct/total*100:.2f}%\")\n#     print(\"-\" * 60)\n#     print(f\"  ✅ True Positive (Bệnh → Bệnh): {true_positive:,}\")\n#     print(f\"  ✅ True Negative (KBệnh → KBệnh): {true_negative:,}\")\n#     print(f\"  ❌ False Positive (KBệnh → Bệnh): {false_positive:,}\")\n#     print(f\"  ❌ False Negative (Bệnh → KBệnh): {false_negative:,}\")\n#     print(\"=\" * 60)\n    \n#     return {\n#         'total_correct': int(correct),\n#         'total_labels': int(total),\n#         'accuracy': correct / total,\n#         'true_positive': int(true_positive),\n#         'true_negative': int(true_negative),\n#         'false_positive': int(false_positive),\n#         'false_negative': int(false_negative)\n#     }\n\n\n# # =================== MAIN ===================\n# if __name__ == \"__main__\":\n#     device = torch.device(DEVICE if torch.cuda.is_available() else \"cpu\")\n#     print(f\"🖥️ Device: {device}\")\n    \n#     # Load model\n#     print(f\"\\n📂 Loading model from: {CHECKPOINT_PATH}\")\n#     model = CSRModel(\n#         num_classes=NUM_CLASSES,\n#         num_prototypes=NUM_PROTOTYPES,\n#         model_name=MODEL_NAME,\n#         pretrained=False\n#     )\n    \n#     ckpt = torch.load(CHECKPOINT_PATH, map_location=device)\n#     if 'model_state_dict' in ckpt:\n#         model.load_state_dict(ckpt['model_state_dict'])\n#     else:\n#         model.load_state_dict(ckpt)\n    \n#     model.to(device)\n#     model.eval()\n#     print(\"✅ Model loaded!\")\n    \n#     # Load test data\n#     image_ids, image_labels = load_test_data(CSV_PATH, IMAGE_DIR)\n    \n#     # Evaluate\n#     all_probs, all_preds, all_labels = evaluate_model(\n#         model, image_ids, image_labels, IMAGE_DIR, device, threshold=THRESHOLD\n#     )\n    \n#     # ===== ĐƠN GIẢN: Đếm số nhãn đúng =====\n#     simple_results = compute_simple_accuracy(all_preds, all_labels)\n    \n#     # ===== CHI TIẾT: Tất cả metrics =====\n#     results = compute_metrics(all_probs, all_preds, all_labels)\n#     print_evaluation_results(results, all_labels)\n    \n#     # Lưu kết quả\n#     print(\"\\n💾 Saving results...\")\n#     np.savez('evaluation_results.npz', \n#              probs=all_probs, \n#              preds=all_preds, \n#              labels=all_labels)\n#     print(\"✅ Done!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T16:02:32.897647Z","iopub.execute_input":"2025-12-10T16:02:32.89818Z","iopub.status.idle":"2025-12-10T16:03:47.601752Z","shell.execute_reply.started":"2025-12-10T16:02:32.898154Z","shell.execute_reply":"2025-12-10T16:03:47.601116Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# \"\"\"\n# Visualize Concept Activation Maps (CAMs) từ model Phase 1\n# \"\"\"\n\n# import torch\n# import torch.nn as nn\n# import torch.nn.functional as F\n# import numpy as np\n# import cv2\n# import matplotlib.pyplot as plt\n# from matplotlib import cm\n# import os\n# import timm\n\n# # =================== CẤU HÌNH ===================\n\n# # Checkpoint và ảnh\n# CHECKPOINT_PATH = \"/kaggle/input/vin-csr-training/checkpoints/csr_phase1.pth\"\n# IMAGE_PATH = \"/kaggle/input/vindr-image-convert/train_png_384/9a5094b2563a1ef3ff50dc5c7ff71345.png\"\n# SAVE_DIR = \"/kaggle/working/cam_outputs\"\n\n# # Model config\n# NUM_CLASSES = 14\n# NUM_PROTOTYPES = 15\n# MODEL_NAME = \"densenet121\"\n# IMG_SIZE = 384\n# DEVICE = \"cuda\"\n\n# # Visualization config\n# TOP_K = 3  # Số CAM hiển thị (top-k theo probability)\n# CAM_THRESHOLD = 0.6  # Ngưỡng để lọc CAM\n# ALPHA = 0.5  # Độ trong suốt overlay\n\n# # Tên các bệnh\n# CLASS_NAMES = [\n#     'Aortic enlargement', 'Atelectasis', 'Calcification', 'Cardiomegaly',\n#     'Consolidation', 'ILD', 'Infiltration', 'Lung Opacity',\n#     'Nodule/Mass', 'Other lesion', 'Pleural effusion', 'Pleural thickening',\n#     'Pneumothorax', 'Pulmonary fibrosis'\n# ]\n\n# # ================================================\n\n\n# class CSRModel(nn.Module):\n#     def __init__(self, num_classes=14, num_prototypes=5, model_name=\"resnet50\", pretrained=True):\n#         super().__init__()\n        \n#         self.backbone = timm.create_model(\n#             model_name, pretrained=pretrained, features_only=True, out_indices=(4,)\n#         )\n#         feature_info = self.backbone.feature_info.get_dicts()[-1]\n#         self.feature_dim = feature_info[\"num_chs\"]\n#         self.concept_head = nn.Conv2d(self.feature_dim, num_classes, kernel_size=1)\n        \n#         self.embedding_dim = 128\n#         self.projector = nn.Sequential(\n#             nn.Linear(self.feature_dim, 512),\n#             nn.ReLU(),\n#             nn.Linear(512, self.embedding_dim)\n#         )\n        \n#         self.prototypes = nn.Parameter(torch.randn(num_classes, num_prototypes, self.embedding_dim))\n#         self.task_head = nn.Linear(num_classes * num_prototypes, num_classes)\n#         self.num_classes = num_classes\n#         self.num_prototypes = num_prototypes\n\n#     def get_features_and_cam(self, x):\n#         if x.size(1) == 1:\n#             x = x.repeat(1, 3, 1, 1)\n#         features = self.backbone(x)[0]\n#         attn_logits = self.concept_head(features)\n#         return features, attn_logits\n    \n#     def forward_phase1(self, x):\n#         \"\"\"Forward cho Phase 1 - dùng GAP trên CAM\"\"\"\n#         _, attn_logits = self.get_features_and_cam(x)\n#         logits = F.adaptive_avg_pool2d(attn_logits, (1, 1)).view(x.size(0), -1)\n#         return {\"logits\": logits, \"attn_maps\": attn_logits}\n\n\n# def load_model(checkpoint_path, device):\n#     \"\"\"Load model từ checkpoint\"\"\"\n#     model = CSRModel(\n#         num_classes=NUM_CLASSES,\n#         num_prototypes=NUM_PROTOTYPES,\n#         model_name=MODEL_NAME,\n#         pretrained=False\n#     )\n    \n#     ckpt = torch.load(checkpoint_path, map_location=device)\n#     if 'model_state_dict' in ckpt:\n#         model.load_state_dict(ckpt['model_state_dict'])\n#     else:\n#         model.load_state_dict(ckpt)\n    \n#     model.to(device)\n#     model.eval()\n#     return model\n\n\n# def preprocess_image(image_path, target_size=384):\n#     \"\"\"Đọc và tiền xử lý ảnh\"\"\"\n#     image = cv2.imread(image_path, cv2.IMREAD_GRAYSCALE)\n#     if image is None:\n#         raise ValueError(f\"Cannot read image: {image_path}\")\n    \n#     original_image = image.copy()\n#     image = cv2.resize(image, (target_size, target_size))\n#     img_norm = image.astype(np.float32) / 255.0\n#     img_tensor = torch.from_numpy(img_norm).unsqueeze(0).unsqueeze(0)\n    \n#     return img_tensor, image, original_image\n\n\n# def get_cam_overlay(original_img, cam, threshold=0.3, alpha=0.5, colormap=cv2.COLORMAP_JET):\n#     \"\"\"\n#     Tạo overlay CAM lên ảnh gốc\n    \n#     Args:\n#         original_img: Ảnh grayscale (H, W)\n#         cam: CAM array (H_cam, W_cam)\n#         threshold: Ngưỡng để lọc CAM\n#         alpha: Độ trong suốt\n#         colormap: OpenCV colormap\n    \n#     Returns:\n#         overlay: Ảnh RGB với CAM overlay\n#     \"\"\"\n#     # Resize CAM về kích thước ảnh gốc\n#     cam_resized = cv2.resize(cam, (original_img.shape[1], original_img.shape[0]))\n    \n#     # Normalize CAM về [0, 1]\n#     cam_min, cam_max = cam_resized.min(), cam_resized.max()\n#     if cam_max - cam_min > 1e-8:\n#         cam_norm = (cam_resized - cam_min) / (cam_max - cam_min)\n#     else:\n#         cam_norm = np.zeros_like(cam_resized)\n    \n#     # Áp dụng threshold\n#     cam_thresholded = np.where(cam_norm > threshold, cam_norm, 0)\n    \n#     # Re-normalize sau threshold\n#     if cam_thresholded.max() > 0:\n#         cam_thresholded = cam_thresholded / cam_thresholded.max()\n    \n#     # Tạo heatmap\n#     heatmap = cv2.applyColorMap(np.uint8(255 * cam_thresholded), colormap)\n#     heatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB)\n    \n#     # Chuyển ảnh gốc sang RGB\n#     img_rgb = cv2.cvtColor(original_img, cv2.COLOR_GRAY2RGB)\n    \n#     # Tạo mask để chỉ overlay vùng có CAM\n#     mask = (cam_thresholded > 0).astype(np.float32)\n#     mask = np.stack([mask] * 3, axis=-1)\n    \n#     # Blend\n#     overlay = img_rgb.astype(np.float32)\n#     overlay = overlay * (1 - mask * alpha) + heatmap.astype(np.float32) * (mask * alpha)\n#     overlay = np.clip(overlay, 0, 255).astype(np.uint8)\n    \n#     return overlay, cam_norm\n\n\n# def visualize_all_cams(model, image_path, device, top_k=5, threshold=0.3, save_path=None):\n#     \"\"\"\n#     Visualize CAM cho tất cả classes (hoặc top-k)\n#     \"\"\"\n#     # Load và preprocess ảnh\n#     img_tensor, resized_img, original_img = preprocess_image(image_path, IMG_SIZE)\n#     img_tensor = img_tensor.to(device)\n    \n#     # Inference\n#     with torch.no_grad():\n#         outputs = model.forward_phase1(img_tensor)\n#         logits = outputs['logits'][0]\n#         cams = outputs['attn_maps'][0]  # [K, H, W]\n#         probs = torch.sigmoid(logits).cpu().numpy()\n    \n#     # Lấy top-k classes theo probability\n#     top_indices = np.argsort(probs)[::-1][:top_k]\n    \n#     # Tạo figure\n#     n_cols = min(3, top_k + 1)\n#     n_rows = (top_k + 1 + n_cols - 1) // n_cols\n#     fig, axes = plt.subplots(n_rows, n_cols, figsize=(5 * n_cols, 5 * n_rows))\n#     axes = axes.flatten() if n_rows > 1 or n_cols > 1 else [axes]\n    \n#     # Ảnh gốc\n#     axes[0].imshow(resized_img, cmap='gray')\n#     axes[0].set_title('Original X-Ray', fontsize=14, fontweight='bold')\n#     axes[0].axis('off')\n    \n#     # Thêm text prediction\n#     pred_text = \"Predictions:\\n\"\n#     for idx in top_indices[:5]:\n#         pred_text += f\"{CLASS_NAMES[idx]}: {probs[idx]*100:.1f}%\\n\"\n#     axes[0].text(0.02, 0.02, pred_text, transform=axes[0].transAxes,\n#                  fontsize=9, verticalalignment='bottom',\n#                  bbox=dict(boxstyle='round', facecolor='white', alpha=0.8))\n    \n#     # CAM cho từng class\n#     for i, idx in enumerate(top_indices):\n#         ax = axes[i + 1]\n        \n#         cam = cams[idx].cpu().numpy()\n#         overlay, cam_norm = get_cam_overlay(resized_img, cam, threshold=threshold)\n        \n#         ax.imshow(overlay)\n#         ax.set_title(f'{CLASS_NAMES[idx]}\\nProb: {probs[idx]*100:.1f}%', fontsize=12)\n#         ax.axis('off')\n        \n#         # Thêm colorbar nhỏ\n#         # coverage = (cam_norm > threshold).sum() / cam_norm.size * 100\n#         # ax.text(0.02, 0.02, f'Coverage: {coverage:.1f}%', transform=ax.transAxes,\n#         #         fontsize=8, bbox=dict(boxstyle='round', facecolor='white', alpha=0.7))\n    \n#     # Ẩn axes thừa\n#     for j in range(len(top_indices) + 1, len(axes)):\n#         axes[j].axis('off')\n    \n#     plt.suptitle(f'Concept Activation Maps (threshold={threshold})', fontsize=16, fontweight='bold')\n#     plt.tight_layout()\n    \n#     if save_path:\n#         os.makedirs(os.path.dirname(save_path), exist_ok=True)\n#         plt.savefig(save_path, dpi=150, bbox_inches='tight')\n#         print(f\"💾 Saved to: {save_path}\")\n    \n#     plt.show()\n    \n#     return probs, cams\n\n\n# def visualize_single_cam(model, image_path, class_idx, device, threshold=0.3, save_path=None):\n#     \"\"\"\n#     Visualize CAM cho 1 class cụ thể với chi tiết hơn\n#     \"\"\"\n#     img_tensor, resized_img, _ = preprocess_image(image_path, IMG_SIZE)\n#     img_tensor = img_tensor.to(device)\n    \n#     with torch.no_grad():\n#         outputs = model.forward_phase1(img_tensor)\n#         logits = outputs['logits'][0]\n#         cams = outputs['attn_maps'][0]\n#         probs = torch.sigmoid(logits).cpu().numpy()\n    \n#     cam = cams[class_idx].cpu().numpy()\n#     prob = probs[class_idx]\n    \n#     # Tạo figure với nhiều views\n#     fig, axes = plt.subplots(2, 3, figsize=(15, 10))\n    \n#     # Row 1: Original, CAM raw, Overlay\n#     axes[0, 0].imshow(resized_img, cmap='gray')\n#     axes[0, 0].set_title('Original', fontsize=12)\n#     axes[0, 0].axis('off')\n    \n#     # CAM raw (resize về kích thước ảnh)\n#     cam_resized = cv2.resize(cam, (resized_img.shape[1], resized_img.shape[0]))\n#     im = axes[0, 1].imshow(cam_resized, cmap='jet')\n#     axes[0, 1].set_title(f'CAM (raw logits)\\nmin={cam.min():.2f}, max={cam.max():.2f}', fontsize=12)\n#     axes[0, 1].axis('off')\n#     plt.colorbar(im, ax=axes[0, 1], fraction=0.046)\n    \n#     # Overlay\n#     overlay, cam_norm = get_cam_overlay(resized_img, cam, threshold=threshold)\n#     axes[0, 2].imshow(overlay)\n#     axes[0, 2].set_title(f'Overlay (threshold={threshold})', fontsize=12)\n#     axes[0, 2].axis('off')\n    \n#     # Row 2: Different thresholds\n#     thresholds = [0.0, 0.3, 0.5]\n#     for i, thresh in enumerate(thresholds):\n#         overlay_t, _ = get_cam_overlay(resized_img, cam, threshold=thresh)\n#         axes[1, i].imshow(overlay_t)\n#         coverage = (cam_norm > thresh).sum() / cam_norm.size * 100\n#         axes[1, i].set_title(f'Threshold={thresh}\\nCoverage: {coverage:.1f}%', fontsize=11)\n#         axes[1, i].axis('off')\n    \n#     plt.suptitle(f'{CLASS_NAMES[class_idx]} - Probability: {prob*100:.1f}%', \n#                  fontsize=16, fontweight='bold')\n#     plt.tight_layout()\n    \n#     if save_path:\n#         os.makedirs(os.path.dirname(save_path), exist_ok=True)\n#         plt.savefig(save_path, dpi=150, bbox_inches='tight')\n#         print(f\"💾 Saved to: {save_path}\")\n    \n#     plt.show()\n\n\n# def visualize_cam_comparison(model, image_path, device, save_path=None):\n#     \"\"\"\n#     So sánh CAM của tất cả 14 classes trong 1 figure\n#     \"\"\"\n#     img_tensor, resized_img, _ = preprocess_image(image_path, IMG_SIZE)\n#     img_tensor = img_tensor.to(device)\n    \n#     with torch.no_grad():\n#         outputs = model.forward_phase1(img_tensor)\n#         logits = outputs['logits'][0]\n#         cams = outputs['attn_maps'][0]\n#         probs = torch.sigmoid(logits).cpu().numpy()\n    \n#     # Figure 4x4 (1 original + 14 CAMs + 1 empty)\n#     fig, axes = plt.subplots(4, 4, figsize=(16, 16))\n#     axes = axes.flatten()\n    \n#     # Original\n#     axes[0].imshow(resized_img, cmap='gray')\n#     axes[0].set_title('Original', fontsize=10, fontweight='bold')\n#     axes[0].axis('off')\n    \n#     # 14 CAMs\n#     for i in range(NUM_CLASSES):\n#         ax = axes[i + 1]\n#         cam = cams[i].cpu().numpy()\n#         overlay, _ = get_cam_overlay(resized_img, cam, threshold=CAM_THRESHOLD)\n        \n#         ax.imshow(overlay)\n#         ax.set_title(f'{CLASS_NAMES[i]}\\n{probs[i]*100:.1f}%', fontsize=9)\n#         ax.axis('off')\n        \n#         # Đánh dấu nếu prob > 50%\n#         if probs[i] > 0.5:\n#             ax.patch.set_edgecolor('red')\n#             ax.patch.set_linewidth(3)\n    \n#     # Ẩn axis cuối\n#     axes[15].axis('off')\n    \n#     plt.suptitle(f'All 14 Concept Activation Maps', fontsize=16, fontweight='bold')\n#     plt.tight_layout()\n    \n#     if save_path:\n#         os.makedirs(os.path.dirname(save_path), exist_ok=True)\n#         plt.savefig(save_path, dpi=150, bbox_inches='tight')\n#         print(f\"💾 Saved to: {save_path}\")\n    \n#     plt.show()\n\n\n# # =================== MAIN ===================\n\n# if __name__ == \"__main__\":\n#     device = torch.device(DEVICE if torch.cuda.is_available() else \"cpu\")\n#     print(f\"🖥️ Device: {device}\")\n    \n#     # Load model\n#     print(f\"\\n📂 Loading model from: {CHECKPOINT_PATH}\")\n#     model = load_model(CHECKPOINT_PATH, device)\n#     print(\"✅ Model loaded!\")\n    \n#     # Tạo output directory\n#     os.makedirs(SAVE_DIR, exist_ok=True)\n#     image_id = os.path.splitext(os.path.basename(IMAGE_PATH))[0]\n    \n#     # 1. Visualize Top-K CAMs\n#     print(f\"\\n🔍 Visualizing Top-{TOP_K} CAMs...\")\n#     visualize_all_cams(\n#         model, IMAGE_PATH, device,\n#         top_k=TOP_K,\n#         threshold=CAM_THRESHOLD,\n#         save_path=os.path.join(SAVE_DIR, f\"{image_id}_top{TOP_K}_cams.png\")\n#     )\n    \n#     # 2. Visualize tất cả 14 CAMs\n#     # print(f\"\\n🔍 Visualizing all 14 CAMs...\")\n#     # visualize_cam_comparison(\n#     #     model, IMAGE_PATH, device,\n#     #     save_path=os.path.join(SAVE_DIR, f\"{image_id}_all_cams.png\")\n#     # )\n    \n#     # 3. Visualize chi tiết 1 class (ví dụ: Cardiomegaly - index 3)\n#     # print(f\"\\n🔍 Visualizing detailed CAM for Cardiomegaly...\")\n#     # visualize_single_cam(\n#     #     model, IMAGE_PATH, class_idx=3, device=device,\n#     #     threshold=CAM_THRESHOLD,\n#     #     save_path=os.path.join(SAVE_DIR, f\"{image_id}_cardiomegaly_detail.png\")\n#     # )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T15:48:46.410309Z","iopub.execute_input":"2025-12-10T15:48:46.410627Z","iopub.status.idle":"2025-12-10T15:48:48.448286Z","shell.execute_reply.started":"2025-12-10T15:48:46.410606Z","shell.execute_reply":"2025-12-10T15:48:48.447485Z"}},"outputs":[],"execution_count":null}]}