{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"datasetVersion","sourceId":13855583,"datasetId":8826439,"databundleVersionId":14617430}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ====================================================\n# CELL 0: KHÔI PHỤC DỮ LIỆU TỪ DATASET VÀO WORKING (TỰ ĐỘNG)\n# ====================================================\nimport os\nimport shutil\n\n# ⚠️ Đường dẫn tới Dataset chứa file backup của bạn. \n# Thường Kaggle sẽ mount theo cấu trúc /kaggle/input/tên-dataset\nINPUT_CACHE_DIR = '/kaggle/input/bonsai-cache-resume' \n\nWORK_DIR = '/kaggle/working'\nOUT_DIR = os.path.join(WORK_DIR, 'model_checkpoints')\n\nprint(\"🔄 Đang khôi phục dữ liệu từ Cache...\")\nos.makedirs(OUT_DIR, exist_ok=True)\n\nif os.path.exists(INPUT_CACHE_DIR):\n    # Quét tự động TẤT CẢ các file có trong thư mục Backup\n    for file_name in os.listdir(INPUT_CACHE_DIR):\n        src = os.path.join(INPUT_CACHE_DIR, file_name)\n        \n        if os.path.isfile(src):\n            # Nếu là file CSV -> Copy ra ngoài thư mục gốc Working\n            if file_name.endswith('.csv'):\n                shutil.copy(src, WORK_DIR)\n                print(f\"✅ Đã copy CSV: {file_name}\")\n            \n            # Nếu là file Model (.pth) hoặc Encoders (.pkl) -> Copy vào model_checkpoints\n            elif file_name.endswith('.pth') or file_name.endswith('.pkl'):\n                shutil.copy(src, OUT_DIR)\n                print(f\"✅ Đã copy Model/Encoder: {file_name}\")\nelse:\n    print(f\"❌ CẢNH BÁO: Không tìm thấy thư mục {INPUT_CACHE_DIR}. Hãy kiểm tra lại tên thư mục Dataset lúc bạn Add Data nhé!\")\n\nprint(\"🎉 Khôi phục hoàn tất! Sẵn sàng chạy.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T03:37:40.861392Z","iopub.execute_input":"2026-08-12T03:37:40.861874Z","iopub.status.idle":"2026-08-12T03:37:44.467964Z","shell.execute_reply.started":"2026-08-12T03:37:40.861842Z","shell.execute_reply":"2026-08-12T03:37:44.467219Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# CELL 1: COMMAND CENTER - IMPORTS & GLOBAL CONFIGURATION\n# ====================================================\n\n# 1. Cài đặt các thư viện bên ngoài (Chạy 1 lần)\n# !pip install -q timm albumentations wandb\n\n# 2. STANDARD LIBRARIES (Hệ thống & Tiện ích)\nimport os\nimport sys\nimport gc\nimport time\nimport math\nimport random\nimport warnings\nimport glob\nfrom pathlib import Path\nfrom tqdm.auto import tqdm # Hiển thị thanh tiến trình đẹp mắt\n\n# 3. DATA & MATH (Xử lý dữ liệu bảng & Toán học)\nimport numpy as np\nimport pandas as pd\n\n# 4. IMAGE PROCESSING (Xử lý ảnh)\nimport cv2\nfrom PIL import Image, ImageFile\nImageFile.LOAD_TRUNCATED_IMAGES = True # Tránh crash khi gặp ảnh bị corrupt một nửa\n\n# 5. VISUALIZATION (Vẽ biểu đồ & Phân tích)\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n# Thiết lập style mặc định cho biểu đồ chuẩn nghiên cứu\nsns.set_theme(style=\"whitegrid\", palette=\"muted\")\nplt.rcParams.update({'font.size': 12})\n\n# 6. MACHINE LEARNING & METRICS (Đánh giá)\nfrom sklearn.model_selection import StratifiedKFold, KFold, train_test_split\nfrom sklearn.metrics import accuracy_score, f1_score, roc_auc_score, confusion_matrix, classification_report\nfrom sklearn.preprocessing import MultiLabelBinarizer, LabelEncoder\n\n# 7. DEEP LEARNING (PyTorch Core)\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR, CosineAnnealingWarmRestarts\nfrom torch.cuda.amp import autocast, GradScaler # Dùng cho Mixed Precision Training\n\n# 8. VISION MODELS & AUGMENTATION\nimport torchvision\nimport torchvision.transforms as T\nimport timm # PyTorch Image Models (Kho pre-trained weights)\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# 9. TRACKING & LOGGING\nimport wandb\n\n# Bỏ qua các cảnh báo gây rối mắt trên Kaggle\nwarnings.filterwarnings('ignore')\n\n# ====================================================\n# OBJECT CẤU HÌNH TRUNG TÂM (CFG)\n# ====================================================\nclass CFG:\n    # --- A. Điều khiển Luồng (Pipeline Control Flags) ---\n    seed = 42\n    debug = False             \n    run_eda = True            \n    check_corrupt_imgs = True \n    use_wandb = False         \n    \n    # --- B. Môi trường & Phần cứng (TRỞ VỀ 1 GPU) ---\n    num_workers = 4           # Trở về 4 luồng đọc mặc định cho 1 GPU\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    pin_memory = False\n    \n    # --- C. Kiến trúc Đường dẫn Kaggle ---\n    root_dir = '/kaggle/input'\n    work_dir = '/kaggle/working'\n    out_dir = os.path.join(work_dir, 'model_checkpoints')\n    master_csv_path = os.path.join(work_dir, 'master_train.csv')\n    \n    # --- D. Dữ liệu & Kích thước ---\n    img_size = 224            \n    n_folds = 5               \n    \n    # Nhãn Đầu ra\n    target_species_col = 'species_label'\n    target_disease_col = 'disease_labels'\n    domain_col = 'domain_source'\n    \n    # --- E. Cấu hình Mạng Backbone ---\n    backbone = 'convnextv2_tiny' \n    pretrained = True\n    drop_rate = 0.2           \n    drop_path_rate = 0.2      \n    \n    # --- F. Huấn luyện (2-Stage Training) ---\n    batch_size = 64           # Giảm lại 64 cho 1 GPU khỏi tràn VRAM\n    epochs_stage1 = 10        \n    epochs_stage2 = 15        \n    \n    # 🔥 ĐÃ GIỮ CÁC CẬP NHẬT CHỐNG NỔ LOSS (NaN)\n    lr = 2e-4                 \n    min_lr = 1e-6             \n    weight_decay = 1e-2\n    \n    # Cấu hình DANN\n    use_dann = True\n    dann_lambda_init = 0.0\n    dann_lambda_max = 1.0\n    \n    # Gradient & Tối ưu RAM\n    gradient_accumulation_steps = 2 \n    max_grad_norm = 2.0       # Giữ mức kẹp 2.0 chống Gradient Explosion\n    amp = True                      \n\n# Tạo thư mục output\nos.makedirs(CFG.out_dir, exist_ok=True)\n\n# ====================================================\n# HÀM HỖ TRỢ HỆ THỐNG\n# ====================================================\ndef seed_everything(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\ndef check_environment():\n    print(\"=\"*50)\n    print(\"📊 BÁO CÁO MÔI TRƯỜNG & PHẦN CỨNG\")\n    print(\"=\"*50)\n    print(f\"Python Version: {sys.version.split()[0]}\")\n    print(f\"PyTorch Version: {torch.__version__}\")\n    if torch.cuda.is_available():\n        print(f\"GPU Name: {torch.cuda.get_device_name(0)}\")\n        print(f\"GPU Count: {torch.cuda.device_count()}\")\n        print(f\"CUDA Version: {torch.version.cuda}\")\n    else:\n        print(\"⚠️ CẢNH BÁO: Không tìm thấy GPU. Đang chạy bằng CPU!\")\n    print(\"=\"*50)\n    print(f\"✅ Đã thiết lập Seed: {CFG.seed}\")\n    print(f\"✅ Thư mục Checkpoint: {CFG.out_dir}\")\n\nseed_everything(CFG.seed)\ncheck_environment()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-08-12T03:37:44.469320Z","iopub.execute_input":"2026-08-12T03:37:44.469572Z","iopub.status.idle":"2026-08-12T03:37:51.819510Z","shell.execute_reply.started":"2026-08-12T03:37:44.469552Z","shell.execute_reply":"2026-08-12T03:37:51.818472Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# CELL 2: EXPLICIT ETL PIPELINE & MASTER SCHEMA BUILDER\n# ====================================================\n\nimport os\nimport pandas as pd\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\n\nprint(\"🚀 Khởi động Trình quản lý Dữ liệu (Data Manager)...\")\n\n# ====================================================\n# 🔄 LOGIC AUTO-SKIP (NHẬN DIỆN DỮ LIỆU ĐÃ LƯU)\n# ====================================================\nif os.path.exists(CFG.master_csv_path) and not CFG.debug:\n    print(f\"⚡ [CACHE HIT] Đã tìm thấy {CFG.master_csv_path}!\")\n    print(\"Bỏ qua bước quét ổ cứng. Đang tải trực tiếp dữ liệu...\")\n    master_df = pd.read_csv(CFG.master_csv_path)\n    \n    print(\"\\n\" + \"=\"*50)\n    print(\"📊 TỔNG QUAN MASTER DATAFRAME (LOADED FROM CACHE)\")\n    print(\"=\"*50)\n    print(f\"Tổng số lượng ảnh: {len(master_df)}\")\n    print(f\"Số lượng Nguồn Domain (Datasets): {master_df['domain_source'].nunique()}\")\n    print(\"=\"*50)\n\nelse:\n    print(\"🔍 Không tìm thấy dữ liệu bộ nhớ đệm. Bắt đầu quá trình ETL quét ổ cứng...\")\n    \n    # ====================================================\n    # A. DANH SÁCH ĐƯỜNG DẪN TĨNH\n    # ====================================================\n    COMPETITION_PATHS = {\n        'pp2020': '/kaggle/input/competitions/plant-pathology-2020-fgvc7',\n        'pp2021': '/kaggle/input/competitions/plant-pathology-2021-fgvc8',\n        'planttraits2024': '/kaggle/input/competitions/planttraits2024'\n    }\n\n    FOLDER_DATASET_PATHS = [\n        '/kaggle/input/datasets/rashidthihan/plant-disease-dataset/plant_disease_dataset/train',\n        '/kaggle/input/datasets/rashikrahmanpritom/plant-disease-recognition-dataset/Train/Train',\n        '/kaggle/input/datasets/saroz014/plant-diseases/train',\n        '/kaggle/input/datasets/mohitsingh1804/plantvillage/PlantVillage/train',\n        '/kaggle/input/datasets/ankursingh12/resized-plant2021/img_sz_256',\n        '/kaggle/input/datasets/ankursingh12/resized-plant2021/img_sz_384',\n        '/kaggle/input/datasets/ankursingh12/resized-plant2021/img_sz_512',\n        '/kaggle/input/datasets/ankursingh12/resized-plant2021/img_sz_640',\n        '/kaggle/input/datasets/riteshranjansaroj/segmented-medicinal-leaf-images/Segmented Medicinal Leaf Images',\n        '/kaggle/input/datasets/vbookshelf/v2-plant-seedlings-dataset',\n        '/kaggle/input/datasets/vipoooool/new-plant-diseases-dataset',\n        '/kaggle/input/datasets/csafrit2/plant-leaves-for-image-classification/Plants_2/train',\n        '/kaggle/input/datasets/tushar5harma/plant-village-dataset-updated/Apple/Train',\n        '/kaggle/input/datasets/tushar5harma/plant-village-dataset-updated/Bell Pepper/Train',\n        '/kaggle/input/datasets/tushar5harma/plant-village-dataset-updated/Cherry/Train',\n        '/kaggle/input/datasets/tushar5harma/plant-village-dataset-updated/Corn (Maize)/Train',\n        '/kaggle/input/datasets/tushar5harma/plant-village-dataset-updated/Grape/Train',\n        '/kaggle/input/datasets/tushar5harma/plant-village-dataset-updated/Peach/Train',\n        '/kaggle/input/datasets/tushar5harma/plant-village-dataset-updated/Potato/Train',\n        '/kaggle/input/datasets/tushar5harma/plant-village-dataset-updated/Strawberry/Train',\n        '/kaggle/input/datasets/tushar5harma/plant-village-dataset-updated/Tomato/Train',\n        '/kaggle/input/datasets/bulentsiyah/plantvillage/PlantVillage_resize_224/PlantVillage_resize_224',\n        '/kaggle/input/datasets/lavaman151/plantifydr-dataset/PlantDiseasesDataset/Apple/train',\n        '/kaggle/input/datasets/lavaman151/plantifydr-dataset/PlantDiseasesDataset/Bell pepper/train',\n        '/kaggle/input/datasets/lavaman151/plantifydr-dataset/PlantDiseasesDataset/Cherry/train',\n        '/kaggle/input/datasets/lavaman151/plantifydr-dataset/PlantDiseasesDataset/Citrus/train',\n        '/kaggle/input/datasets/lavaman151/plantifydr-dataset/PlantDiseasesDataset/Corn/train',\n        '/kaggle/input/datasets/lavaman151/plantifydr-dataset/PlantDiseasesDataset/Grape/train',\n        '/kaggle/input/datasets/lavaman151/plantifydr-dataset/PlantDiseasesDataset/Peach/train',\n        '/kaggle/input/datasets/lavaman151/plantifydr-dataset/PlantDiseasesDataset/Potato/train',\n        '/kaggle/input/datasets/lavaman151/plantifydr-dataset/PlantDiseasesDataset/Strawberry/train',\n        '/kaggle/input/datasets/lavaman151/plantifydr-dataset/PlantDiseasesDataset/Tomato/train',\n        '/kaggle/input/datasets/abdallahalidev/plantvillage-dataset/color',\n        '/kaggle/input/datasets/abdallahalidev/plantvillage-dataset/grayscale',\n        '/kaggle/input/datasets/abdallahalidev/plantvillage-dataset/segmented',\n        '/kaggle/input/datasets/yudhaislamisulistya/plants-type-datasets/split_ttv_dataset_type_of_plants/Train_Set_Folder'\n    ]\n\n    # ====================================================\n    # B. HÀM SUY LUẬN NHÃN\n    # ====================================================\n    def extract_labels_from_foldername(folder_name):\n        name = folder_name.lower().strip()\n        if '___' in name:\n            species, disease = name.split('___', 1)\n            disease = 'healthy' if 'healthy' in disease else disease\n            return species.strip(), [disease.strip()]\n        if ' healthy' in name:\n            species = name.replace(' healthy', '').split('(')[0].strip()\n            return species, ['healthy']\n        if ' diseased' in name:\n            species = name.replace(' diseased', '').split('(')[0].strip()\n            return species, ['diseased']\n        disease_keywords = ['rot', 'rust', 'scab', 'mildew', 'blight', 'spot', 'virus', 'powdery', 'healthy']\n        if any(kw in name for kw in disease_keywords) and len(name.split('_')) <= 3:\n            return 'unknown_species', ['healthy' if 'healthy' in name else name]\n        return name, ['unknown_disease']\n\n    # ====================================================\n    # C. CÁC HÀM PARSERS\n    # ====================================================\n    def parse_competitions():\n        data = []\n        pp20_path = COMPETITION_PATHS['pp2020']\n        if os.path.exists(os.path.join(pp20_path, 'train.csv')):\n            df = pd.read_csv(os.path.join(pp20_path, 'train.csv'))\n            img_dir = os.path.join(pp20_path, 'images')\n            disease_cols = [c for c in df.columns if c != 'image_id']\n            for _, row in df.iterrows():\n                img_name = row['image_id'] if str(row['image_id']).endswith('.jpg') else str(row['image_id']) + '.jpg'\n                active = [c for c in disease_cols if row[c] == 1]\n                data.append({'image_path': os.path.join(img_dir, img_name), 'domain_source': 'pp2020', 'species_label': 'apple', 'disease_labels': active})\n\n        pp21_path = COMPETITION_PATHS['pp2021']\n        if os.path.exists(os.path.join(pp21_path, 'train.csv')):\n            df = pd.read_csv(os.path.join(pp21_path, 'train.csv'))\n            img_dir = os.path.join(pp21_path, 'train_images')\n            for _, row in df.iterrows():\n                diseases = row['labels'].lower().split(' ')\n                data.append({'image_path': os.path.join(img_dir, str(row['image'])), 'domain_source': 'pp2021', 'species_label': 'apple', 'disease_labels': diseases})\n                \n        pt24_path = COMPETITION_PATHS['planttraits2024']\n        if os.path.exists(os.path.join(pt24_path, 'train.csv')):\n            df = pd.read_csv(os.path.join(pt24_path, 'train.csv'))\n            img_dir = os.path.join(pt24_path, 'train_images')\n            for _, row in df.iterrows():\n                img_name = str(row['id']) + '.jpeg'\n                data.append({'image_path': os.path.join(img_dir, img_name), 'domain_source': 'planttraits24', 'species_label': 'unknown_species', 'disease_labels': ['unknown_disease']})\n\n        return pd.DataFrame(data)\n\n    def parse_explicit_folders():\n        data = []\n        seen_filenames = set()\n        img_exts = {'.jpg', '.jpeg', '.png', '.JPG', '.PNG'}\n        \n        for base_path in tqdm(FOLDER_DATASET_PATHS, desc=\"Quét thư mục Dataset\"):\n            path_obj = Path(base_path)\n            if not path_obj.exists():\n                continue\n                \n            domain_name = path_obj.parts[4] if len(path_obj.parts) > 4 else path_obj.name\n            \n            for ext in img_exts:\n                for img_path in path_obj.rglob(f'*{ext}'):\n                    if 'test' in img_path.parts or 'val' in img_path.parts or 'valid' in img_path.parts:\n                        continue\n                    filename = img_path.name\n                    if filename in seen_filenames:\n                        continue\n                    seen_filenames.add(filename)\n                    folder_name = img_path.parent.name\n                    species, diseases = extract_labels_from_foldername(folder_name)\n                    data.append({\n                        'image_path': str(img_path),\n                        'domain_source': domain_name,\n                        'species_label': species,\n                        'disease_labels': diseases\n                    })\n        return pd.DataFrame(data)\n\n    # ====================================================\n    # D. KÍCH HOẠT VÀ GỘP DỮ LIỆU\n    # ====================================================\n    print(\"1/2: Xử lý các cuộc thi Kaggle (CSV)...\")\n    df_comps = parse_competitions()\n\n    print(\"\\n2/2: Xử lý các Dataset ảnh dạng Thư mục...\")\n    df_folders = parse_explicit_folders()\n\n    all_dfs = []\n    if not df_comps.empty: all_dfs.append(df_comps)\n    if not df_folders.empty: all_dfs.append(df_folders)\n\n    if all_dfs:\n        master_df = pd.concat(all_dfs, ignore_index=True)\n        \n        if CFG.debug:\n            master_df = master_df.sample(n=min(500, len(master_df)), random_state=CFG.seed).reset_index(drop=True)\n            print(\"🐞 CHẾ ĐỘ DEBUG: Chỉ sử dụng tập dữ liệu thu nhỏ.\")\n\n        print(\"\\n🛡️ Đang kiểm tra tính hợp lệ của đường dẫn file vật lý...\")\n        master_df['is_valid'] = master_df['image_path'].apply(lambda x: os.path.exists(x))\n        master_df = master_df[master_df['is_valid']].drop(columns=['is_valid']).reset_index(drop=True)\n\n        master_df.to_csv(CFG.master_csv_path, index=False)\n        \n        print(\"\\n\" + \"=\"*70)\n        print(\"📊 TỔNG QUAN MASTER DATAFRAME (EXPLICIT PATHS)\")\n        print(\"=\"*70)\n        print(f\"Tổng số lượng ảnh hợp lệ: {len(master_df)}\")\n        print(f\"Số lượng Nguồn Domain (Datasets): {master_df['domain_source'].nunique()}\")\n        print(\"=\"*70)\n    else:\n        print(\"❌ LỖI NGHIÊM TRỌNG: Quá trình quét trả về 0 ảnh.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T03:37:51.820601Z","iopub.execute_input":"2026-08-12T03:37:51.821390Z","iopub.status.idle":"2026-08-12T03:37:52.723612Z","shell.execute_reply.started":"2026-08-12T03:37:51.821362Z","shell.execute_reply":"2026-08-12T03:37:52.722486Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# CELL 3: DATA CLEANING, LABEL ENCODING & EDA\n# ====================================================\n\nimport os\nimport ast\nimport joblib\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom concurrent.futures import ThreadPoolExecutor\nfrom sklearn.preprocessing import MultiLabelBinarizer, LabelEncoder\n\nprint(\"🔍 Bắt đầu quá trình Data Cleaning & EDA...\")\n\ncleaned_csv_path = os.path.join(CFG.work_dir, 'cleaned_train.csv')\nenc_species_path = os.path.join(CFG.out_dir, 'species_encoder.pkl')\nenc_disease_path = os.path.join(CFG.out_dir, 'disease_mlb.pkl')\n\n# ====================================================\n# HÀM SỬA LỖI ĐỌC MẢNG NUMPY TỪ CSV (FIX SYNTAX ERROR)\n# ====================================================\ndef parse_array_string(x):\n    \"\"\"Xử lý lỗi thiếu dấu phẩy khi đọc mảng Numpy dạng chuỗi từ file CSV.\"\"\"\n    if isinstance(x, str):\n        # Trừng hợp 1: Chuỗi Numpy không có dấu phẩy (VD: \"[0 0 1 0]\")\n        if '[' in x and ',' not in x:\n            clean_str = x.replace('[', '').replace(']', '').replace('\\n', ' ')\n            return [int(float(i)) for i in clean_str.split()]\n        # Trường hợp 2: Python list chuẩn có dấu phẩy (VD: \"[0, 0, 1, 0]\")\n        return ast.literal_eval(x)\n    return x\n\n# ====================================================\n# 🔄 LOGIC AUTO-SKIP (NHẬN DIỆN DỮ LIỆU ĐÃ CLEAN)\n# ====================================================\nif os.path.exists(cleaned_csv_path) and os.path.exists(enc_species_path) and os.path.exists(enc_disease_path) and not CFG.debug:\n    print(\"⚡ [CACHE HIT] Tìm thấy dữ liệu đã Clean & Encode. Đang tải trực tiếp...\")\n    df = pd.read_csv(cleaned_csv_path)\n    \n    # Khôi phục định dạng Python List từ file CSV\n    df['disease_labels'] = df['disease_labels'].apply(lambda x: ast.literal_eval(x) if isinstance(x, str) else x)\n    \n    # 🔥 ĐÃ ÁP DỤNG HÀM SỬA LỖI VÀO ĐÂY\n    if 'disease_target' in df.columns:\n        df['disease_target'] = df['disease_target'].apply(parse_array_string)\n    \n    # Tải lại bộ Encoders\n    species_le = joblib.load(enc_species_path)\n    mlb = joblib.load(enc_disease_path)\n    \n    CFG.species_classes = len(species_le.classes_)\n    CFG.disease_classes = len(mlb.classes_)\n    \n    print(f\"✅ Đã khôi phục thành công! (Loài: {CFG.species_classes}, Bệnh: {CFG.disease_classes})\")\n\nelse:\n    print(\"🧹 Không tìm thấy Cache. Bắt đầu Dọn dẹp dữ liệu và Encode nhãn...\")\n    \n    # 1. TẢI DỮ LIỆU TỪ BƯỚC TRƯỚC\n    df = pd.read_csv(CFG.master_csv_path)\n    df['disease_labels'] = df['disease_labels'].apply(lambda x: ast.literal_eval(x) if isinstance(x, str) else x)\n\n    # ====================================================\n    # A. LÀM SẠCH DỮ LIỆU (DATA CLEANING)\n    # ====================================================\n    def verify_image(img_path):\n        \"\"\"Mở thử ảnh bằng cv2 để kiểm tra file có bị hỏng (corrupt) không\"\"\"\n        try:\n            if not os.path.exists(img_path): return False\n            img = cv2.imread(img_path)\n            if img is None: return False\n            return True\n        except:\n            return False\n\n    if CFG.check_corrupt_imgs:\n        print(f\"Bắt đầu quét {len(df)} ảnh để tìm file lỗi. Quá trình này có thể mất vài phút...\")\n        with ThreadPoolExecutor(max_workers=CFG.num_workers * 2) as executor:\n            valid_mask = list(tqdm(executor.map(verify_image, df['image_path']), total=len(df), desc=\"Verifying Images\"))\n        \n        invalid_count = len(df) - sum(valid_mask)\n        df = df[valid_mask].reset_index(drop=True)\n        print(f\"🗑️ Đã xóa {invalid_count} ảnh lỗi/không tồn tại. Dữ liệu hợp lệ còn lại: {len(df)} ảnh.\")\n\n    # ====================================================\n    # B. MÃ HÓA NHÃN (LABEL ENCODING)\n    # ====================================================\n    species_le = LabelEncoder()\n    df['species_target'] = species_le.fit_transform(df['species_label'])\n    joblib.dump(species_le, enc_species_path)\n\n    mlb = MultiLabelBinarizer()\n    disease_encoded = mlb.fit_transform(df['disease_labels'])\n    joblib.dump(mlb, enc_disease_path)\n\n    # 🔥 FIX BẬC HAI: Ép Numpy Array thành dạng List chuẩn trước khi lưu để tránh sinh ra lỗi sau này\n    df['disease_target'] = [list(arr) for arr in disease_encoded]\n\n    CFG.species_classes = len(species_le.classes_)\n    CFG.disease_classes = len(mlb.classes_)\n\n    print(\"\\n\" + \"=\"*50)\n    print(\"📌 THÔNG TIN LABEL ENCODING\")\n    print(\"=\"*50)\n    print(f\"- Số lượng loài cây: {CFG.species_classes}\")\n    print(f\"- Số lượng bệnh lý: {CFG.disease_classes}\")\n\n    df.to_csv(cleaned_csv_path, index=False)\n    print(f\"✅ Đã lưu tập dữ liệu sạch vào {cleaned_csv_path}\")\n\n# ====================================================\n# C. EXPLORATORY DATA ANALYSIS (EDA) & DOMAIN GAP\n# ====================================================\nif CFG.run_eda:\n    print(\"\\n📊 Đang tạo biểu đồ phân tích...\")\n    \n    fig, axes = plt.subplots(1, 2, figsize=(20, 6))\n    \n    # Biểu đồ 1\n    sns.countplot(y='species_label', data=df, ax=axes[0], order=df['species_label'].value_counts().index[:15])\n    axes[0].set_title('Top 15 Loài Cây phổ biến nhất')\n    \n    # Biểu đồ 2: Đảm bảo encode lại để vẽ biểu đồ nếu đang dùng Cache\n    if 'mlb' not in locals():\n        mlb = joblib.load(enc_disease_path)\n    \n    disease_encoded = mlb.transform(df['disease_labels'])\n    disease_counts = pd.DataFrame(disease_encoded, columns=mlb.classes_).sum().sort_values(ascending=False).head(15)\n    sns.barplot(x=disease_counts.values, y=disease_counts.index, ax=axes[1])\n    axes[1].set_title('Top 15 Bệnh Lý phổ biến nhất')\n    \n    plt.tight_layout()\n    plt.show()\n    \n    # Trực quan hóa DOMAIN GAP\n    print(\"\\n🖼️ TRỰC QUAN HÓA DOMAIN GAP\")\n    domains = df['domain_source'].unique()\n    \n    if len(domains) >= 2:\n        plot_domains = domains[:6]\n        fig, axes = plt.subplots(len(plot_domains), 4, figsize=(16, 4 * len(plot_domains)))\n        if len(plot_domains) == 1: axes = [axes] # Fix lỗi matrix nếu chỉ có 1 row\n        \n        fig.suptitle('So sánh sự khác biệt về Background giữa các Nguồn dữ liệu (Tối đa 6 nguồn)', fontsize=16)\n        \n        for i, domain in enumerate(plot_domains):\n            domain_df = df[df['domain_source'] == domain]\n            n_samples = min(4, len(domain_df))\n            sample_df = domain_df.sample(n=n_samples, random_state=CFG.seed)\n            \n            for j in range(4):\n                if j < n_samples:\n                    row = sample_df.iloc[j]\n                    img = cv2.imread(row['image_path'])\n                    if img is not None:\n                        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n                        img = cv2.resize(img, (224, 224))\n                        axes[i][j].imshow(img)\n                        disease_str = str(row['disease_labels'])[:30] + \"...\" if len(str(row['disease_labels'])) > 30 else str(row['disease_labels'])\n                        axes[i][j].set_title(f\"[{domain}]\\n{row['species_label']} - {disease_str}\", fontsize=9)\n                axes[i][j].axis('off')\n                \n        plt.tight_layout()\n        plt.show()\n    else:\n        print(\"💡 Chỉ có 1 Domain, bỏ qua phần so sánh Domain Gap.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T03:37:52.725440Z","iopub.execute_input":"2026-08-12T03:37:52.725812Z","iopub.status.idle":"2026-08-12T03:38:08.449480Z","shell.execute_reply.started":"2026-08-12T03:37:52.725788Z","shell.execute_reply":"2026-08-12T03:38:08.448686Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# CELL 4: AUGMENTATION & PREPROCESSING ENGINE\n# ====================================================\n\nprint(\"🎨 Đang khởi tạo bộ máy Augmentation (Albumentations)...\")\n\ndef get_train_transforms():\n    \"\"\"\n    Pipeline biến đổi dành cho tập Huấn luyện.\n    Kết hợp Spatial (Không gian) và Pixel-level (Điểm ảnh) transforms.\n    \"\"\"\n    return A.Compose([\n        # 1. Resize với padding để giữ nguyên tỷ lệ khung hình (Aspect Ratio)\n        # Hữu ích cho ảnh Bonsai vì dáng cây không bị bóp méo\n        A.LongestMaxSize(max_size=CFG.img_size, p=1.0),\n        A.PadIfNeeded(min_height=CFG.img_size, min_width=CFG.img_size, \n                      border_mode=cv2.BORDER_CONSTANT, value=[0, 0, 0], p=1.0),\n        \n        # 2. Biến đổi hình học cơ bản (An toàn cho lá cây)\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.15, rotate_limit=45, p=0.6),\n        \n        # 3. Giả lập môi trường thực tế (Pixel-level để chống Domain Gap)\n        A.OneOf([\n            A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2, p=1.0),\n            A.HueSaturationValue(hue_shift_limit=20, sat_shift_limit=30, val_shift_limit=20, p=1.0),\n            A.RandomGamma(p=1.0),\n        ], p=0.7),\n        \n        # 4. Giả lập nhiễu từ camera điện thoại (Blur, Noise)\n        A.OneOf([\n            A.GaussianBlur(blur_limit=(3, 5), p=1.0),\n            A.MotionBlur(blur_limit=5, p=1.0),\n            A.GaussNoise(var_limit=(10.0, 50.0), p=1.0),\n        ], p=0.4),\n        \n        # 5. Giả lập thời tiết & Bóng râm ngoài vườn\n        A.OneOf([\n            A.RandomShadow(p=1.0),\n            A.RandomFog(p=1.0),\n        ], p=0.2), # Xác suất thấp để không làm hỏng quá nhiều data\n        \n        # 6. Che khuất ngẫu nhiên (CoarseDropout / CutOut)\n        # Ép mô hình học đặc trưng phân tán, không phụ thuộc vào 1 điểm duy nhất\n        A.CoarseDropout(max_holes=8, max_height=int(CFG.img_size * 0.1), \n                        max_width=int(CFG.img_size * 0.1), min_holes=1, \n                        fill_value=0, p=0.5),\n        \n        # 7. Chuẩn hóa & Đưa về Tensor PyTorch\n        A.Normalize(\n            mean=[0.485, 0.456, 0.406], # Mean chuẩn của ImageNet\n            std=[0.229, 0.224, 0.225],  # Std chuẩn của ImageNet\n            max_pixel_value=255.0, \n            p=1.0\n        ),\n        ToTensorV2(p=1.0),\n    ])\n\ndef get_valid_transforms():\n    \"\"\"\n    Pipeline biến đổi dành cho tập Validation / Test.\n    Tuyệt đối không dùng random augmentation, chỉ Resize và Normalize.\n    \"\"\"\n    return A.Compose([\n        A.LongestMaxSize(max_size=CFG.img_size, p=1.0),\n        A.PadIfNeeded(min_height=CFG.img_size, min_width=CFG.img_size, \n                      border_mode=cv2.BORDER_CONSTANT, value=[0, 0, 0], p=1.0),\n        A.Normalize(\n            mean=[0.485, 0.456, 0.406],\n            std=[0.229, 0.224, 0.225],\n            max_pixel_value=255.0,\n            p=1.0\n        ),\n        ToTensorV2(p=1.0),\n    ])\n\n# ====================================================\n# KIỂM TRA TRỰC QUAN AUGMENTATION (Visual Check)\n# ====================================================\n\nif CFG.run_eda:\n    # Lấy thử 1 đường dẫn ảnh ngẫu nhiên từ df (Giả sử df đã được định nghĩa ở Cell 3)\n    try:\n        sample_img_path = df['image_path'].iloc[random.randint(0, len(df)-1)]\n        image = cv2.imread(sample_img_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        train_transforms = get_train_transforms()\n        \n        # Vẽ biểu đồ 1 ảnh gốc và 4 biến thể\n        fig, axes = plt.subplots(1, 5, figsize=(20, 5))\n        axes[0].imshow(image)\n        axes[0].set_title('Ảnh Gốc')\n        axes[0].axis('off')\n        \n        for i in range(1, 5):\n            # Áp dụng Augmentation (Lưu ý: Phải tháo Normalize & ToTensor để matplotlib vẽ được)\n            # Ở đây chúng ta tạm bypass Normalize để trực quan hóa\n            aug_pipeline = A.Compose(train_transforms.transforms[:-2]) \n            augmented = aug_pipeline(image=image)['image']\n            \n            axes[i].imshow(augmented)\n            axes[i].set_title(f'Augmented {i}')\n            axes[i].axis('off')\n            \n        plt.suptitle(\"Kiểm tra Tác động của Augmentation\", fontsize=16)\n        plt.tight_layout()\n        plt.show()\n    except Exception as e:\n        print(f\"Bỏ qua bước vẽ hình Augmentation: {e}\")\n\nprint(\"✅ Đã khởi tạo hoàn tất Pipeline Tiền xử lý ảnh.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T03:38:08.451203Z","iopub.execute_input":"2026-08-12T03:38:08.451637Z","iopub.status.idle":"2026-08-12T03:38:09.011843Z","shell.execute_reply.started":"2026-08-12T03:38:08.451590Z","shell.execute_reply":"2026-08-12T03:38:09.010837Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# CELL 5: DATASET CLASSES & DATALOADERS\n# ====================================================\n\nimport joblib\nfrom sklearn.model_selection import StratifiedKFold\nfrom torch.utils.data import Dataset, DataLoader\n\nprint(\"🗂️ Đang khởi tạo Dataset và Dataloaders...\")\n\nenc_domain_path = os.path.join(CFG.out_dir, 'domain_encoder.pkl')\nfolds_path = os.path.join(CFG.work_dir, 'folds.csv')\n\n# ====================================================\n# 1. MÃ HÓA NHÃN DOMAIN (Nhận diện Cache)\n# ====================================================\nif os.path.exists(enc_domain_path) and not CFG.debug:\n    print(\"⚡ [CACHE HIT] Tìm thấy Domain Encoder đã lưu. Đang tải...\")\n    domain_le = joblib.load(enc_domain_path)\n    # df['domain_target'] đã được tạo sẵn trong cleaned_train.csv nếu ta lưu chuẩn\n    # Nếu chưa có, ta transform lại\n    if 'domain_target' not in df.columns:\n        df['domain_target'] = domain_le.transform(df['domain_source'])\nelse:\n    print(\"Mã hóa Domain Encoder mới...\")\n    domain_le = LabelEncoder()\n    df['domain_target'] = domain_le.fit_transform(df['domain_source'])\n    joblib.dump(domain_le, enc_domain_path)\n\nCFG.domain_classes = len(domain_le.classes_)\n\n# ====================================================\n# A. CROSS-VALIDATION SPLIT (Nhận diện Cache)\n# ====================================================\nif 'fold' in df.columns and not CFG.debug:\n    print(\"⚡ [CACHE HIT] DataFrame đã chứa cột Folds. Bỏ qua bước Stratified K-Fold.\")\nelse:\n    print(\"Tiến hành phân chia Folds mới...\")\n    # Tạo khóa phân tầng kết hợp (Domain + Species)\n    df['stratify_key'] = df['domain_source'] + \"_\" + df['species_target'].astype(str)\n\n    skf = StratifiedKFold(n_splits=CFG.n_folds, shuffle=True, random_state=CFG.seed)\n    df['fold'] = -1\n\n    for fold, (train_idx, val_idx) in enumerate(skf.split(X=df, y=df['stratify_key'])):\n        df.loc[val_idx, 'fold'] = fold\n        \n    df.to_csv(folds_path, index=False)\n    print(f\"✅ Đã chia {CFG.n_folds} Folds và lưu vào {folds_path}\")\n\nif CFG.run_eda:\n    # Kiểm tra nhanh xem Fold 0 có phân phối Domain tốt không\n    print(\"\\nPhân phối Domain trong Fold 0 (Validation Set):\")\n    print(df[df['fold'] == 0]['domain_source'].value_counts(normalize=True) * 100)\n\n# ====================================================\n# B. CUSTOM PYTORCH DATASET\n# ====================================================\nclass BonsaiMultiHeadDataset(Dataset):\n    def __init__(self, df, transforms=None):\n        self.df = df\n        self.file_names = df['image_path'].values\n        self.species_labels = df['species_target'].values\n        self.disease_labels = np.array(df['disease_target'].tolist()) # Chuyển Multi-hot list thành Numpy Array\n        self.domain_labels = df['domain_target'].values\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        # 1. Đọc ảnh bằng OpenCV\n        img_path = self.file_names[index]\n        image = cv2.imread(img_path)\n        \n        # Xử lý an toàn nếu ảnh bị lỗi (Dù đã clean ở Cell 3, nhưng đề phòng đứt kết nối disk)\n        if image is None:\n            image = np.zeros((CFG.img_size, CFG.img_size, 3), dtype=np.uint8)\n        else:\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n            \n        # 2. Áp dụng Augmentation\n        if self.transforms:\n            res = self.transforms(image=image)\n            image = res['image'] # Trả về PyTorch Tensor (C, H, W)\n            \n        # 3. Trích xuất nhãn\n        species = torch.tensor(self.species_labels[index], dtype=torch.long)\n        disease = torch.tensor(self.disease_labels[index], dtype=torch.float32) # Sigmoid/BCE cần Float32\n        domain = torch.tensor(self.domain_labels[index], dtype=torch.long)\n        \n        # 4. Trả về Tuple\n        return image, species, disease, domain\n\n# ====================================================\n# C. DATALOADER GENERATOR\n# ====================================================\ndef prepare_loaders(df, fold):\n    \"\"\"\n    Tạo Dataloader cho quá trình huấn luyện theo từng Fold.\n    \"\"\"\n    train_df = df[df.fold != fold].reset_index(drop=True)\n    valid_df = df[df.fold == fold].reset_index(drop=True)\n    \n    train_dataset = BonsaiMultiHeadDataset(train_df, transforms=get_train_transforms())\n    valid_dataset = BonsaiMultiHeadDataset(valid_df, transforms=get_valid_transforms())\n    \n    train_loader = DataLoader(\n        train_dataset, \n        batch_size=CFG.batch_size, \n        shuffle=True, \n        num_workers=CFG.num_workers, \n        pin_memory=CFG.pin_memory, \n        drop_last=True # Rất quan trọng: Bỏ qua batch cuối nếu bị lẻ để tránh lỗi BatchNorm\n    )\n    \n    valid_loader = DataLoader(\n        valid_dataset, \n        batch_size=CFG.batch_size * 2, # Inference tốn ít RAM hơn Train, nhân 2 để chạy nhanh hơn\n        shuffle=False, \n        num_workers=CFG.num_workers, \n        pin_memory=CFG.pin_memory, \n        drop_last=False\n    )\n    \n    return train_loader, valid_loader\n\n# Kiểm thử nhanh Dataloader\ntry:\n    temp_train_loader, _ = prepare_loaders(df, fold=0)\n    temp_images, temp_species, temp_diseases, temp_domains = next(iter(temp_train_loader))\n    \n    print(\"\\n📦 KIỂM TRA TENSOR TỪ DATALOADER:\")\n    print(f\"- Image Batch Shape: {temp_images.shape}\")\n    print(f\"- Species Labels Shape: {temp_species.shape} | Type: {temp_species.dtype}\")\n    print(f\"- Disease Labels Shape: {temp_diseases.shape} | Type: {temp_diseases.dtype}\")\n    print(f\"- Domain Labels Shape: {temp_domains.shape} | Type: {temp_domains.dtype}\")\n    print(\"✅ Dataloader hoạt động hoàn hảo!\")\nexcept Exception as e:\n    print(f\"❌ LỖI DATALOADER: {e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T03:38:09.013083Z","iopub.execute_input":"2026-08-12T03:38:09.013422Z","iopub.status.idle":"2026-08-12T03:38:20.540907Z","shell.execute_reply.started":"2026-08-12T03:38:09.013399Z","shell.execute_reply":"2026-08-12T03:38:20.539867Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# CELL 6: MODEL ARCHITECTURE & DANN IMPLEMENTATION\n# ====================================================\n\nprint(\"🧠 Đang xây dựng Kiến trúc Mạng Neural (Multi-Head + DANN)...\")\n\n# ====================================================\n# A. GRADIENT REVERSAL LAYER (GRL)\n# ====================================================\nclass GradientReversalFn(torch.autograd.Function):\n    @staticmethod\n    def forward(ctx, x, alpha):\n        # Lưu lại hằng số alpha (lambda) để dùng trong backward\n        ctx.alpha = alpha\n        \n        # Forward pass: Dữ liệu đi qua không thay đổi\n        return x.view_as(x)\n\n    @staticmethod\n    def backward(ctx, grad_output):\n        # Backward pass: Nhân gradient với âm alpha\n        output = grad_output.neg() * ctx.alpha\n        \n        # Trả về gradient cho x, trả về None cho alpha (vì alpha là hằng số)\n        return output, None\n\ndef grad_reverse(x, alpha=1.0):\n    return GradientReversalFn.apply(x, alpha)\n\n# ====================================================\n# B. MULTI-HEAD BONSAI MODEL\n# ====================================================\nclass BonsaiMultiHeadModel(nn.Module):\n    def __init__(self, backbone_name=CFG.backbone, pretrained=CFG.pretrained):\n        super().__init__()\n        \n        # 1. KHỞI TẠO BACKBONE TỪ TIMM\n        # num_classes=0: Loại bỏ lớp Fully Connected cuối cùng\n        # global_pool='avg': Sử dụng Global Average Pooling để tạo Vector 1D\n        self.backbone = timm.create_model(\n            backbone_name, \n            pretrained=pretrained, \n            num_classes=0, \n            global_pool='avg',\n            drop_rate=CFG.drop_rate,\n            drop_path_rate=CFG.drop_path_rate\n        )\n        \n        # Lấy kích thước vector đầu ra của Backbone (VD: ConvNeXt-Tiny ra vector 768 chiều)\n        self.in_features = self.backbone.num_features\n        \n        # 2. KHỞI TẠO CÁC NHÁNH ĐẦU RA (HEADS)\n        \n        # Head 1: Phân loại Loài cây (Multi-Class)\n        self.species_head = nn.Sequential(\n            nn.Linear(self.in_features, 512),\n            nn.BatchNorm1d(512),\n            nn.SiLU(), # SiLU (Swish) tốt hơn ReLU cho mạng hiện đại\n            nn.Dropout(0.3),\n            nn.Linear(512, CFG.species_classes)\n        )\n        \n        # Head 2: Phân loại Bệnh lý (Multi-Label)\n        # Bệnh lý phức tạp hơn, cần Head sâu hơn một chút\n        self.disease_head = nn.Sequential(\n            nn.Linear(self.in_features, 512),\n            nn.BatchNorm1d(512),\n            nn.SiLU(),\n            nn.Dropout(0.4),\n            nn.Linear(512, CFG.disease_classes)\n        )\n        \n        # Head 3: Phân loại Nguồn Domain (Dành cho DANN)\n        self.domain_head = nn.Sequential(\n            nn.Linear(self.in_features, 256),\n            nn.BatchNorm1d(256),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(256, CFG.domain_classes)\n        )\n\n    def forward(self, x, alpha=0.0):\n        \"\"\"\n        Quá trình truyền xuôi của mạng.\n        Args:\n            x: Tensor ảnh (B, C, H, W)\n            alpha: Trọng số của GRL (lambda). Nếu alpha = 0, DANN bị vô hiệu hóa.\n        \"\"\"\n        # 1. Trích xuất đặc trưng (Feature Extraction)\n        # Kết quả trả về: Vector (Batch_size, in_features)\n        features = self.backbone(x)\n        \n        # 2. Dự đoán Loài cây và Bệnh lý\n        species_logits = self.species_head(features)\n        disease_logits = self.disease_head(features)\n        \n        # 3. Kích hoạt DANN và Dự đoán Domain\n        if CFG.use_dann:\n            # Áp dụng Gradient Reversal Layer lên vector đặc trưng\n            reversed_features = grad_reverse(features, alpha)\n            domain_logits = self.domain_head(reversed_features)\n        else:\n            # Nếu không dùng DANN, tạo dummy tensor để tránh lỗi unpack\n            domain_logits = torch.zeros((x.size(0), CFG.domain_classes)).to(features.device)\n            \n        return species_logits, disease_logits, domain_logits\n\n# ====================================================\n# KIỂM THỬ MÔ HÌNH (Sanity Check)\n# ====================================================\ntry:\n    # Khởi tạo mô hình\n    model = BonsaiMultiHeadModel()\n    model.to(CFG.device)\n    \n    # Tạo một Batch ảnh giả định\n    dummy_images = torch.randn(CFG.batch_size, 3, CFG.img_size, CFG.img_size).to(CFG.device)\n    dummy_alpha = 0.5\n    \n    # Chạy thử (Forward pass)\n    spec_out, dis_out, dom_out = model(dummy_images, dummy_alpha)\n    \n    print(\"\\n🏗️ KIỂM TRA KIẾN TRÚC MẠNG TẠI ĐẦU RA:\")\n    print(f\"- Backbone Model: {CFG.backbone}\")\n    print(f\"- Species Output Shape: {spec_out.shape} (Kỳ vọng: Batch, Species_Classes)\")\n    print(f\"- Disease Output Shape: {dis_out.shape} (Kỳ vọng: Batch, Disease_Classes)\")\n    print(f\"- Domain Output Shape: {dom_out.shape} (Kỳ vọng: Batch, Domain_Classes)\")\n    print(\"✅ Kiến trúc Model Multi-Head DANN hoạt động chính xác!\")\n    \n    # Giải phóng RAM/VRAM để tiết kiệm cho các Cell sau\n    del model, dummy_images, spec_out, dis_out, dom_out\n    torch.cuda.empty_cache()\n    gc.collect()\n\nexcept Exception as e:\n    print(f\"❌ LỖI KHỞI TẠO MÔ HÌNH: {e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T03:38:20.542585Z","iopub.execute_input":"2026-08-12T03:38:20.543332Z","iopub.status.idle":"2026-08-12T03:38:25.978029Z","shell.execute_reply.started":"2026-08-12T03:38:20.543302Z","shell.execute_reply":"2026-08-12T03:38:25.977234Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# CELL 7: LOSS FUNCTIONS & METRICS FORMULATION\n# ====================================================\n\nprint(\"⚖️ Đang khởi tạo Hệ thống Hàm Mất mát (Loss) và Đo lường (Metrics)...\")\n\n# ====================================================\n# A. CUSTOM FOCAL LOSS CHO NHÁNH BỆNH LÝ (MULTI-LABEL)\n# ====================================================\nclass MultiLabelFocalLoss(nn.Module):\n    def __init__(self, alpha=0.25, gamma=2.0, reduction='mean'):\n        \"\"\"\n        Focal Loss được thiết kế chuyên biệt cho Multi-Label Classification.\n        Giúp mô hình tập trung học các bệnh hiếm gặp (minority classes).\n        \"\"\"\n        super(MultiLabelFocalLoss, self).__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.reduction = reduction\n        self.bce_with_logits = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, inputs, targets):\n        # 1. Tính BCE Loss ban đầu (chưa giảm chiều)\n        bce_loss = self.bce_with_logits(inputs, targets)\n        \n        # 2. Tính xác suất dự đoán (Sigmoid)\n        probs = torch.sigmoid(inputs)\n        \n        # 3. Tính p_t (Xác suất dự đoán đúng class)\n        p_t = probs * targets + (1 - probs) * (1 - targets)\n        \n        # 4. Tính hệ số alpha_t\n        alpha_t = self.alpha * targets + (1 - self.alpha) * (1 - targets)\n        \n        # 5. Áp dụng công thức Focal Loss\n        focal_weight = alpha_t * (1 - p_t) ** self.gamma\n        focal_loss = focal_weight * bce_loss\n        \n        if self.reduction == 'mean':\n            return focal_loss.mean()\n        elif self.reduction == 'sum':\n            return focal_loss.sum()\n        return focal_loss\n\n# ====================================================\n# B. MULTI-HEAD LOSS COMPOSER\n# ====================================================\nclass BonsaiMultiTaskLoss(nn.Module):\n    def __init__(self):\n        super(BonsaiMultiTaskLoss, self).__init__()\n        \n        # 1. Nhánh Phân loại Loài cây (Multi-class: 1 ảnh = 1 cây duy nhất)\n        self.species_criterion = nn.CrossEntropyLoss()\n        \n        # 2. Nhánh Phân loại Bệnh (Multi-label: 1 cây = N bệnh cùng lúc)\n        self.disease_criterion = MultiLabelFocalLoss(alpha=0.25, gamma=2.0)\n        \n        # 3. Nhánh Phân biệt Nguồn Dữ liệu (Domain)\n        self.domain_criterion = nn.CrossEntropyLoss()\n\n    def forward(self, preds, targets):\n        # Unpack Data\n        species_logits, disease_logits, domain_logits = preds\n        species_targets, disease_targets, domain_targets = targets\n        \n        # Tính toán từng Loss độc lập\n        l_species = self.species_criterion(species_logits, species_targets)\n        l_disease = self.disease_criterion(disease_logits, disease_targets)\n        \n        # Chỉ tính Domain Loss nếu mô hình đang kích hoạt DANN (Có dự đoán Domain)\n        if CFG.use_dann:\n            l_domain = self.domain_criterion(domain_logits, domain_targets)\n        else:\n            l_domain = torch.tensor(0.0).to(l_species.device)\n            \n        # Tổng hợp Loss (Chú ý: Gradient ngược đã được xử lý bằng GRL ở Cell 6)\n        total_loss = l_species + l_disease + l_domain\n        \n        # Lưu trữ các giá trị thành phần để visualize\n        loss_dict = {\n            'total_loss': total_loss.item(),\n            'species_loss': l_species.item(),\n            'disease_loss': l_disease.item(),\n            'domain_loss': l_domain.item() if CFG.use_dann else 0.0\n        }\n        \n        return total_loss, loss_dict\n\n# ====================================================\n# C. BỘ MÁY ĐO LƯỜNG CHỈ SỐ (METRICS EVALUATOR)\n# ====================================================\nclass MetricEvaluator:\n    def __init__(self):\n        # Reset các biến tích lũy sau mỗi Epoch\n        self.reset()\n        \n    def reset(self):\n        self.species_preds = []\n        self.species_trues = []\n        \n        self.disease_preds = []\n        self.disease_trues = []\n        \n        self.domain_preds = []\n        self.domain_trues = []\n        \n    def update(self, preds, targets):\n        \"\"\"Lưu trữ dự đoán của từng Batch vào bộ nhớ tạm để tính tổng ở cuối Epoch\"\"\"\n        species_logits, disease_logits, domain_logits = preds\n        species_targets, disease_targets, domain_targets = targets\n        \n        # Xử lý Nhánh Species (Argmax)\n        self.species_preds.extend(torch.argmax(species_logits, dim=1).cpu().numpy())\n        self.species_trues.extend(species_targets.cpu().numpy())\n        \n        # Xử lý Nhánh Disease (Sigmoid + Threshold 0.5)\n        dis_probs = torch.sigmoid(disease_logits).cpu().detach().numpy()\n        self.disease_preds.extend((dis_probs > 0.5).astype(int))\n        self.disease_trues.extend(disease_targets.cpu().numpy())\n        \n        # Xử lý Nhánh Domain\n        if CFG.use_dann:\n            self.domain_preds.extend(torch.argmax(domain_logits, dim=1).cpu().numpy())\n            self.domain_trues.extend(domain_targets.cpu().numpy())\n            \n    def compute(self):\n        \"\"\"Tính toán và trả về các chỉ số cuối cùng\"\"\"\n        metrics = {}\n        \n        # 1. Điểm Accuracy cho Species\n        metrics['species_acc'] = accuracy_score(self.species_trues, self.species_preds)\n        \n        # 2. Điểm F1-Score (Macro) cho Disease\n        # Macro F1 rất quan trọng vì nó trừng phạt các mô hình bỏ qua bệnh hiếm\n        metrics['disease_f1_macro'] = f1_score(self.disease_trues, self.disease_preds, average='macro', zero_division=0)\n        metrics['disease_f1_micro'] = f1_score(self.disease_trues, self.disease_preds, average='micro', zero_division=0)\n        \n        # 3. Điểm Accuracy cho Domain (Nếu càng thấp càng tốt vì chứng tỏ DANN đã lừa được mô hình)\n        if CFG.use_dann and len(self.domain_trues) > 0:\n            metrics['domain_acc'] = accuracy_score(self.domain_trues, self.domain_preds)\n        else:\n            metrics['domain_acc'] = 0.0\n            \n        return metrics\n\n# Khởi tạo thử nghiệm\ncriterion = BonsaiMultiTaskLoss().to(CFG.device)\nevaluator = MetricEvaluator()\nprint(\"✅ Hoàn tất khởi tạo Loss & Metrics. Hệ thống đã sẵn sàng cho Training Loop.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T03:38:25.979321Z","iopub.execute_input":"2026-08-12T03:38:25.980108Z","iopub.status.idle":"2026-08-12T03:38:25.995234Z","shell.execute_reply.started":"2026-08-12T03:38:25.980084Z","shell.execute_reply":"2026-08-12T03:38:25.994224Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# CELL 8: THE TRAINING ENGINE (AUTO-RESUME CAPABILITY)\n# ====================================================\n\nimport torch\nimport torch.nn as nn\nfrom torch.cuda.amp import autocast, GradScaler\nimport numpy as np\nimport os\nimport gc\nfrom tqdm.auto import tqdm\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\nprint(\"🚀 Đang khởi động Động cơ Huấn luyện (1 GPU - Smart Checkpoint)...\")\n\ndef get_dann_alpha(current_epoch, max_epochs):\n    p = current_epoch / max_epochs\n    alpha = (2.0 / (1.0 + np.exp(-10 * p))) - 1.0\n    return float(alpha)\n\ndef train_one_epoch(model, dataloader, optimizer, criterion, scaler, epoch, stage, max_epochs):\n    model.train()\n    evaluator.reset()\n    \n    total_loss_epoch = 0\n    progress_bar = tqdm(enumerate(dataloader), total=len(dataloader), desc=f\"Train Stage {stage} - Epoch {epoch+1}\")\n    \n    alpha = 0.0 if stage == 1 else get_dann_alpha(epoch, max_epochs)\n    \n    optimizer.zero_grad()\n    \n    for step, (images, species, diseases, domains) in progress_bar:\n        images = images.to(CFG.device, non_blocking=True)\n        species = species.to(CFG.device, non_blocking=True)\n        diseases = diseases.to(CFG.device, non_blocking=True)\n        domains = domains.to(CFG.device, non_blocking=True)\n        targets = (species, diseases, domains)\n        \n        with autocast(enabled=CFG.amp):\n            preds = model(images, alpha=alpha)\n            loss, loss_dict = criterion(preds, targets)\n            loss = loss / CFG.gradient_accumulation_steps\n            \n        scaler.scale(loss).backward()\n        \n        if (step + 1) % CFG.gradient_accumulation_steps == 0 or (step + 1) == len(dataloader):\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.max_grad_norm)\n            \n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n            \n        evaluator.update(preds, targets)\n        total_loss_epoch += loss.item() * CFG.gradient_accumulation_steps\n        \n        progress_bar.set_postfix(\n            loss=f\"{loss_dict['total_loss']:.4f}\", \n            alpha=f\"{alpha:.3f}\" if stage == 2 else \"0.0\"\n        )\n        \n    metrics = evaluator.compute()\n    avg_loss = total_loss_epoch / len(dataloader)\n    return avg_loss, metrics\n\n@torch.no_grad()\ndef valid_one_epoch(model, dataloader, criterion):\n    model.eval()\n    evaluator.reset()\n    \n    total_loss_epoch = 0\n    progress_bar = tqdm(enumerate(dataloader), total=len(dataloader), desc=\"Validating\")\n    \n    for step, (images, species, diseases, domains) in progress_bar:\n        images = images.to(CFG.device, non_blocking=True)\n        species = species.to(CFG.device, non_blocking=True)\n        diseases = diseases.to(CFG.device, non_blocking=True)\n        domains = domains.to(CFG.device, non_blocking=True)\n        targets = (species, diseases, domains)\n        \n        preds = model(images, alpha=0.0)\n        loss, loss_dict = criterion(preds, targets)\n        \n        evaluator.update(preds, targets)\n        total_loss_epoch += loss.item()\n        \n    metrics = evaluator.compute()\n    avg_loss = total_loss_epoch / len(dataloader)\n    return avg_loss, metrics\n\ndef run_training(fold=0):\n    print(f\"\\n{'='*60}\\n🔥 BẮT ĐẦU HUẤN LUYỆN - FOLD {fold}/{CFG.n_folds - 1}\\n{'='*60}\")\n    \n    train_loader, valid_loader = prepare_loaders(df, fold)\n    \n    model = BonsaiMultiHeadModel().to(CFG.device)\n    criterion = BonsaiMultiTaskLoss().to(CFG.device)\n    scaler = GradScaler(enabled=CFG.amp)\n    \n    # ====================================================\n    # 🔄 LOGIC AUTO-RESUME (XỬ LÝ ĐƯỢC CẢ FILE CŨ & MỚI)\n    # ====================================================\n    best_score = 0.0\n    start_stage = 1\n    start_epoch = 0\n    checkpoint = None\n    \n    latest_path = os.path.join(CFG.out_dir, f'latest_model_fold_{fold}.pth')\n    if os.path.exists(latest_path):\n        print(f\"🔄 ĐÃ TÌM THẤY BẢN LƯU! Đang khôi phục tiến độ từ: {latest_path}\")\n        checkpoint = torch.load(latest_path, map_location=CFG.device)\n        \n        # Kế thừa dữ liệu thông minh\n        if isinstance(checkpoint, dict) and 'model_state_dict' in checkpoint:\n            # 1. NẾU LÀ ĐỊNH DẠNG MỚI CỦA AUTO-RESUME\n            model.load_state_dict(checkpoint['model_state_dict'])\n            best_score = checkpoint.get('best_score', 0.0)\n            start_stage = checkpoint.get('stage', 1)\n            start_epoch = checkpoint.get('epoch', -1) + 1\n            print(f\"✅ Sẵn sàng chạy tiếp: Đang ở Stage {start_stage} - Bắt đầu từ Epoch {start_epoch + 1}\")\n        else:\n            # 2. NẾU LÀ ĐỊNH DẠNG CŨ CỦA BẠN CHỈ LƯU WEIGHTS\n            model.load_state_dict(checkpoint)\n            print(\"⚠️ Phát hiện bản lưu định dạng cũ (chỉ chứa Trọng số mạng). Sẽ chạy tiếp mô hình với kiến thức hiện có nhưng bắt đầu từ Stage 1 - Epoch 1.\")\n            checkpoint = None # Đặt lại None để Optimizer không bị lỗi\n    else:\n        print(\"✨ Không tìm thấy bản lưu cũ. Bắt đầu huấn luyện từ số 0.\")\n\n    # ----------------------------------------------------\n    # GIAI ĐOẠN 1: WARMUP\n    # ----------------------------------------------------\n    if start_stage == 1:\n        print(f\"\\n[STAGE 1] WARMUP FOLD {fold}: Học đặc trưng phân loại cơ bản...\")\n        optimizer_s1 = AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\n        scheduler_s1 = CosineAnnealingLR(optimizer_s1, T_max=CFG.epochs_stage1, eta_min=CFG.min_lr)\n        \n        # Load lại trạng thái Optimizer nếu đang train dở (chỉ cho định dạng mới)\n        if checkpoint and 'optimizer_state_dict' in checkpoint:\n            optimizer_s1.load_state_dict(checkpoint['optimizer_state_dict'])\n            scheduler_s1.load_state_dict(checkpoint['scheduler_state_dict'])\n        \n        for epoch in range(start_epoch, CFG.epochs_stage1):\n            train_loss, train_metrics = train_one_epoch(model, train_loader, optimizer_s1, criterion, scaler, epoch, stage=1, max_epochs=CFG.epochs_stage1)\n            valid_loss, valid_metrics = valid_one_epoch(model, valid_loader, criterion)\n            scheduler_s1.step()\n            \n            print(f\"S1-Epoch {epoch+1} | Val Loss: {valid_loss:.4f} | Spc Acc: {valid_metrics['species_acc']:.4f} | Dis F1: {valid_metrics['disease_f1_macro']:.4f}\")\n            \n            # 🔥 TỪ NAY VỀ SAU HỆ THỐNG SẼ LƯU \"GÓI DỮ LIỆU TỔNG HỢP\" MỚI\n            torch.save({\n                'stage': 1,\n                'epoch': epoch,\n                'best_score': best_score,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer_s1.state_dict(),\n                'scheduler_state_dict': scheduler_s1.state_dict()\n            }, latest_path)\n            \n        start_epoch = 0 # Hoàn tất Stage 1, reset về 0 để sang Stage 2\n\n    # ----------------------------------------------------\n    # GIAI ĐOẠN 2: DOMAIN ALIGNMENT\n    # ----------------------------------------------------\n    if start_stage <= 2:\n        print(f\"\\n[STAGE 2] DOMAIN ALIGNMENT FOLD {fold}: Kích hoạt DANN và Fine-tuning...\")\n        optimizer_s2 = AdamW(model.parameters(), lr=CFG.lr * 0.1, weight_decay=CFG.weight_decay)\n        scheduler_s2 = CosineAnnealingLR(optimizer_s2, T_max=CFG.epochs_stage2, eta_min=CFG.min_lr)\n        \n        # Load lại trạng thái Optimizer nếu đang train Stage 2\n        if start_stage == 2 and checkpoint and 'optimizer_state_dict' in checkpoint:\n            optimizer_s2.load_state_dict(checkpoint['optimizer_state_dict'])\n            scheduler_s2.load_state_dict(checkpoint['scheduler_state_dict'])\n        \n        for epoch in range(start_epoch, CFG.epochs_stage2):\n            train_loss, train_metrics = train_one_epoch(model, train_loader, optimizer_s2, criterion, scaler, epoch, stage=2, max_epochs=CFG.epochs_stage2)\n            valid_loss, valid_metrics = valid_one_epoch(model, valid_loader, criterion)\n            scheduler_s2.step()\n            \n            print(f\"S2-Epoch {epoch+1} | Val Loss: {valid_loss:.4f} | Spc Acc: {valid_metrics['species_acc']:.4f} | Dis F1: {valid_metrics['disease_f1_macro']:.4f}\")\n            \n            # 🔥 LƯU \"GÓI DỮ LIỆU TỔNG HỢP\" CHO STAGE 2\n            torch.save({\n                'stage': 2,\n                'epoch': epoch,\n                'best_score': best_score,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer_s2.state_dict(),\n                'scheduler_state_dict': scheduler_s2.state_dict()\n            }, latest_path)\n            \n            # Lưu Best Model thuần túy (Chỉ lấy weights đi Inference)\n            current_score = valid_metrics['disease_f1_macro']\n            if current_score > best_score:\n                best_score = current_score\n                save_path = os.path.join(CFG.out_dir, f'best_model_fold_{fold}.pth')\n                torch.save(model.state_dict(), save_path)\n                print(f\"⭐ [MỚI] Đã cập nhật Checkpoint XUẤT SẮC NHẤT (F1: {best_score:.4f})\")\n                \n            gc.collect()\n            torch.cuda.empty_cache()\n    \n    # Dọn dẹp GPU trước khi qua Fold mới\n    del model, train_loader, valid_loader\n    if 'optimizer_s1' in locals(): del optimizer_s1, scheduler_s1\n    if 'optimizer_s2' in locals(): del optimizer_s2, scheduler_s2\n    gc.collect()\n    torch.cuda.empty_cache()\n\n# ====================================================\n# THỰC THI TOÀN BỘ CÁC FOLD\n# ====================================================\nif __name__ == '__main__':\n    # SỬA Ở ĐÂY: Thay range(CFG.n_folds) thành range(1, CFG.n_folds) \n    # Để nó BỎ QUA Fold 0 và bắt đầu chạy thẳng từ Fold 1\n    for fold_idx in range(1, CFG.n_folds):\n        run_training(fold=fold_idx)\n        \n    print(\"\\n🎉 HOÀN TẤT HUẤN LUYỆN TOÀN BỘ HỆ THỐNG!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T03:38:25.996287Z","iopub.execute_input":"2026-08-12T03:38:25.996642Z","execution_failed":"2026-08-12T11:58:04.980Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# CELL 9: ERROR ANALYSIS & HARD EXAMPLE MINING\n# ====================================================\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.metrics import confusion_matrix, classification_report\nimport joblib\n\nprint(\"🕵️ Đang khởi động module Phân tích Lỗi (Error Analysis)...\")\n\ndef load_best_model(fold=0):\n    model = BonsaiMultiHeadModel().to(CFG.device)\n    model_path = os.path.join(CFG.out_dir, f'best_model_fold_{fold}.pth')\n    \n    if os.path.exists(model_path):\n        model.load_state_dict(torch.load(model_path, map_location=CFG.device))\n        print(f\"✅ Đã tải thành công trọng số: {model_path}\")\n    else:\n        print(f\"⚠️ Cảnh báo: Không tìm thấy file {model_path}. Sẽ dùng model khởi tạo ngẫu nhiên để test code.\")\n    \n    model.eval()\n    return model\n\n@torch.no_grad()\ndef run_error_analysis(fold=0):\n    model = load_best_model(fold)\n    \n    # 1. Khởi tạo Pipeline biến đổi dành riêng cho Inference (Hình ảnh sắc nét hơn)\n    # Tự động Scale up kích thước ảnh test lên x1.5 lần so với lúc train để tăng độ chính xác\n    test_img_size = int(CFG.img_size * 1.5) \n    print(f\"🖼️ Đang cấu hình Inference Resolution: {test_img_size}x{test_img_size}\")\n    \n    inference_transforms = A.Compose([\n        A.LongestMaxSize(max_size=test_img_size, p=1.0),\n        A.PadIfNeeded(min_height=test_img_size, min_width=test_img_size, \n                      border_mode=cv2.BORDER_CONSTANT, value=[0, 0, 0], p=1.0),\n        A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225], max_pixel_value=255.0, p=1.0),\n        ToTensorV2(p=1.0),\n    ])\n    \n    # 2. Tái tạo lại Validation Dataloader\n    valid_df = df[df.fold == fold].reset_index(drop=True)\n    valid_dataset = BonsaiMultiHeadDataset(valid_df, transforms=inference_transforms)\n    # Inference cần ít RAM hơn Train, ta nhân đôi Batch Size để chạy nhanh hơn\n    valid_loader = DataLoader(valid_dataset, batch_size=CFG.batch_size * 2, shuffle=False, num_workers=CFG.num_workers)\n    \n    all_species_trues = []\n    all_species_preds = []\n    all_disease_trues = []\n    all_disease_probs = []\n    \n    print(\"⏳ Đang chạy suy luận (Inference) trên toàn bộ tập Validation...\")\n    for images, species, diseases, domains in tqdm(valid_loader):\n        images = images.to(CFG.device)\n        \n        # Forward pass với AMP (Mixed Precision) giúp tăng tốc độ Inference lên 2 lần\n        with autocast(enabled=CFG.amp):\n            spec_logits, dis_logits, _ = model(images, alpha=0.0)\n        \n        # Nhánh Species (Argmax)\n        all_species_preds.extend(torch.argmax(spec_logits, dim=1).cpu().numpy())\n        all_species_trues.extend(species.cpu().numpy())\n        \n        # Nhánh Disease (Sigmoid)\n        dis_probs = torch.sigmoid(dis_logits).cpu().numpy()\n        all_disease_probs.extend(dis_probs)\n        all_disease_trues.extend(diseases.cpu().numpy())\n        \n    all_disease_trues = np.array(all_disease_trues)\n    all_disease_probs = np.array(all_disease_probs)\n    all_disease_preds = (all_disease_probs > 0.5).astype(int)\n\n    # ====================================================\n    # A. CONFUSION MATRIX (LOÀI CÂY)\n    # ====================================================\n    if CFG.run_eda:\n        print(\"\\n📊 1. MA TRẬN NHẦM LẪN (LOÀI CÂY)\")\n        cm = confusion_matrix(all_species_trues, all_species_preds)\n        plt.figure(figsize=(10, 8))\n        \n        species_le = joblib.load(os.path.join(CFG.out_dir, 'species_encoder.pkl'))\n        class_names = species_le.classes_\n        \n        sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=class_names, yticklabels=class_names)\n        plt.title('Confusion Matrix - Phân loại Loài cây', fontsize=15)\n        plt.xlabel('Dự đoán của AI')\n        plt.ylabel('Nhãn Thực tế')\n        plt.xticks(rotation=45)\n        plt.tight_layout()\n        plt.show()\n\n    # ====================================================\n    # B. CLASSIFICATION REPORT (BỆNH LÝ)\n    # ====================================================\n    print(\"\\n📋 2. BÁO CÁO CHI TIẾT BỆNH LÝ (MULTI-LABEL)\")\n    mlb = joblib.load(os.path.join(CFG.out_dir, 'disease_mlb.pkl'))\n    print(classification_report(all_disease_trues, all_disease_preds, target_names=mlb.classes_, zero_division=0))\n\n    # ====================================================\n    # C. HARD EXAMPLE MINING (TRUY VẾT ẢNH DỰ ĐOÁN SAI)\n    # ====================================================\n    print(\"\\n🔍 3. KHAI THÁC CÁC CA DỰ ĐOÁN SAI LỆCH NHẤT (HARD EXAMPLES)\")\n    \n    # Tính sai số (Mean Squared Error) giữa xác suất và thực tế\n    errors = np.mean(np.square(all_disease_probs - all_disease_trues), axis=1)\n    \n    # Lấy ra index của Top 8 ảnh có sai số lớn nhất\n    top_wrong_idx = np.argsort(errors)[-8:][::-1]\n    \n    if CFG.run_eda:\n        fig, axes = plt.subplots(2, 4, figsize=(20, 10))\n        axes = axes.flatten()\n        \n        for i, idx in enumerate(top_wrong_idx):\n            img_path = valid_df.iloc[idx]['image_path']\n            true_species = species_le.inverse_transform([all_species_trues[idx]])[0]\n            pred_species = species_le.inverse_transform([all_species_preds[idx]])[0]\n            \n            true_dis_idx = np.where(all_disease_trues[idx] == 1)[0]\n            pred_dis_idx = np.where(all_disease_preds[idx] == 1)[0]\n            true_dis = mlb.classes_[true_dis_idx] if len(true_dis_idx) > 0 else ['healthy']\n            pred_dis = mlb.classes_[pred_dis_idx] if len(pred_dis_idx) > 0 else ['healthy']\n            \n            img = cv2.imread(img_path)\n            if img is not None:\n                img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n                # Plot ảnh gốc không resize để quan sát rõ nhất lỗi của AI\n                axes[i].imshow(img)\n            \n            title_color = 'red' if list(true_dis) != list(pred_dis) else 'green'\n            title = f\"Thực: {true_dis}\\nĐoán: {pred_dis}\\n---\\nCây: {true_species}\\nAI: {pred_species}\"\n            \n            axes[i].set_title(title, color=title_color, fontsize=10)\n            axes[i].axis('off')\n            \n        plt.suptitle(\"Top 8 Hình Ảnh Bị AI Phân Tích Sai Nặng Nhất\", fontsize=18, fontweight='bold', color='darkred')\n        plt.tight_layout()\n        plt.show()\n\n# Thực thi Module\nif __name__ == '__main__':\n    run_error_analysis(fold=0)","metadata":{"trusted":true,"execution":{"execution_failed":"2026-08-12T11:58:04.980Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# CELL 10: INFERENCE, TTA & PRODUCTION EXPORT\n# ====================================================\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nprint(\"🚀 Đang khởi động module Triển khai (Inference & Export)...\")\n\n# ====================================================\n# 1. TẢI MÔ HÌNH VÀ BỘ TỪ ĐIỂN (ENCODERS)\n# ====================================================\ndef load_inference_components(fold=0):\n    # Tải Model\n    model = BonsaiMultiHeadModel().to(CFG.device)\n    model.load_state_dict(torch.load(os.path.join(CFG.out_dir, f'best_model_fold_{fold}.pth'), map_location=CFG.device))\n    model.eval() # Bắt buộc chuyển sang chế độ đánh giá\n    \n    # Tải Encoders\n    species_le = joblib.load(os.path.join(CFG.out_dir, 'species_encoder.pkl'))\n    disease_mlb = joblib.load(os.path.join(CFG.out_dir, 'disease_mlb.pkl'))\n    \n    return model, species_le, disease_mlb\n\n# ====================================================\n# 2. HÀM DỰ ĐOÁN VỚI TEST TIME AUGMENTATION (TTA)\n# ====================================================\n@torch.no_grad()\ndef predict_single_image(image_path, model, species_le, disease_mlb, threshold=0.5):\n    \"\"\"\n    Đọc một ảnh, áp dụng TTA và trả về kết quả định dạng chuẩn.\n    \"\"\"\n    # 1. Tiền xử lý ảnh\n    image = cv2.imread(image_path)\n    if image is None: return None\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    \n    # 🔥 NÂNG CẤP DYNAMIC RESOLUTION (Giống Cell 9)\n    # Tăng độ phân giải lúc Test để mô hình nhìn rõ bệnh hơn\n    test_img_size = int(CFG.img_size * 1.5)\n    \n    inference_transforms = A.Compose([\n        A.LongestMaxSize(max_size=test_img_size, p=1.0),\n        A.PadIfNeeded(min_height=test_img_size, min_width=test_img_size, \n                      border_mode=cv2.BORDER_CONSTANT, value=[0, 0, 0], p=1.0),\n        A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225], max_pixel_value=255.0, p=1.0),\n        ToTensorV2(p=1.0),\n    ])\n    \n    tensor_img = inference_transforms(image=image)['image'].unsqueeze(0).to(CFG.device) # Shape: (1, C, H, W)\n    \n    # 2. Áp dụng TTA (Test Time Augmentation) bằng các phép toán Tensor\n    img_orig = tensor_img\n    img_hflip = torch.flip(tensor_img, dims=[3]) # Lật theo chiều Width\n    img_vflip = torch.flip(tensor_img, dims=[2]) # Lật theo chiều Height\n    \n    batch_tta = torch.cat([img_orig, img_hflip, img_vflip], dim=0) # Shape: (3, C, H, W)\n    \n    # 3. Chạy Suy luận (Inference) với AMP để tăng tốc\n    with autocast(enabled=CFG.amp):\n        spec_logits, dis_logits, _ = model(batch_tta, alpha=0.0)\n    \n    # 4. Tính xác suất và Lấy trung bình cộng của 3 phiên bản TTA\n    spec_probs = F.softmax(spec_logits, dim=1).mean(dim=0) \n    dis_probs = torch.sigmoid(dis_logits).mean(dim=0)\n    \n    # 5. Dịch ngược kết quả (Decoding)\n    top_species_idx = torch.argmax(spec_probs).item()\n    species_name = species_le.inverse_transform([top_species_idx])[0]\n    species_confidence = spec_probs[top_species_idx].item() * 100\n    \n    dis_probs_np = dis_probs.cpu().numpy()\n    active_disease_indices = np.where(dis_probs_np > threshold)[0]\n    \n    if len(active_disease_indices) > 0:\n        disease_names = disease_mlb.classes_[active_disease_indices].tolist()\n        disease_confidence = np.mean(dis_probs_np[active_disease_indices]) * 100\n    else:\n        disease_names = ['healthy (Không phát hiện bệnh)']\n        disease_confidence = (1.0 - np.max(dis_probs_np)) * 100\n\n    return {\n        'File': os.path.basename(image_path),\n        'Tên Cây': species_name.capitalize(),\n        'Bệnh Lý': \", \".join(disease_names),\n        'Độ tin cậy Cây (%)': f\"{species_confidence:.2f}%\",\n        'Độ tin cậy Bệnh (%)': f\"{disease_confidence:.2f}%\"\n    }\n\n# ====================================================\n# 3. XUẤT MÔ HÌNH (PRODUCTION EXPORT)\n# ====================================================\ndef export_to_production(model, dummy_input_shape=(1, 3, 336, 336)):\n    \"\"\"Xuất mô hình sang định dạng TorchScript để chạy trên server C++ hoặc Mobile.\"\"\"\n    model.eval()\n    dummy_input = torch.randn(dummy_input_shape).to(CFG.device)\n    dummy_alpha = torch.tensor(0.0).to(CFG.device) # Fix an toàn cho kiểu dữ liệu của hàm forward\n    \n    try:\n        # TorchScript yêu cầu mọi input đều phải là Tensor (kể cả tham số alpha)\n        # strict=False giúp lờ đi các dictionary return rườm rà (nếu có)\n        traced_model = torch.jit.trace(model, (dummy_input, dummy_alpha), strict=False)\n        export_path = os.path.join(CFG.work_dir, 'bonsai_multihead_production.pt')\n        traced_model.save(export_path)\n        print(f\"📦 Đã đóng gói TorchScript thành công: {export_path}\")\n    except Exception as e:\n        print(f\"⚠️ Không thể export TorchScript: {e}\")\n\n# ====================================================\n# 4. CHẠY THỬ NGHIỆM THỰC TẾ\n# ====================================================\nif __name__ == '__main__':\n    try:\n        prod_model, prod_species_le, prod_disease_mlb = load_inference_components(fold=0)\n        \n        # Lấy an toàn 5 ảnh nếu DF tồn tại\n        if 'df' in globals() and not df.empty:\n            sample_images = df[df.fold == 0]['image_path'].sample(5, random_state=CFG.seed).values\n            \n            results = []\n            for img_path in sample_images:\n                res = predict_single_image(img_path, prod_model, prod_species_le, prod_disease_mlb)\n                if res: results.append(res)\n                \n            print(\"\\n\" + \"=\"*80)\n            print(\"🎯 KẾT QUẢ DỰ ĐOÁN TỪ HỆ THỐNG AI (Inference Output)\")\n            print(\"=\"*80)\n            results_df = pd.DataFrame(results)\n            display(results_df) \n            print(\"=\"*80)\n        else:\n            print(\"💡 Không tìm thấy DataFrame test, bỏ qua dự đoán thử nghiệm.\")\n        \n        # Xuất file Production với độ phân giải cao đã được định nghĩa ở trên (x1.5 = 336)\n        export_to_production(prod_model, dummy_input_shape=(1, 3, int(CFG.img_size*1.5), int(CFG.img_size*1.5)))\n        \n    except Exception as e:\n        print(f\"⚠️ Bỏ qua Test do lỗi hoặc chưa có file Weights: {e}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2026-08-12T11:58:04.981Z"}},"outputs":[],"execution_count":null}]}