{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"2559deda","cell_type":"code","source":"!pip install -q timm","metadata":{"vscode":{"languageId":"plaintext"}},"outputs":[],"execution_count":null},{"id":"25bf6c9a","cell_type":"code","source":"import os\nimport copy\nimport pandas as pd\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom torchvision.transforms import InterpolationMode\nfrom PIL import Image\nimport timm\nfrom sklearn.model_selection import train_test_split\nimport matplotlib.pyplot as plt\nfrom tqdm.notebook import tqdm\nimport csv","metadata":{},"outputs":[],"execution_count":null},{"id":"ee0eab70","cell_type":"code","source":"# ==========================================\n# 1. CẤU HÌNH THÔNG SỐ & ĐƯỜNG DẪN\n# ==========================================\nIMG_DIR = \"/kaggle/input/datasets/trankimhuu/images-datasets-of-big2015/train_images/kaggle/working/train_images\" \nCSV_PATH = \"/kaggle/input/datasets/trankimhuu/images-datasets-of-big2015/trainLabels.csv\"\nMODEL_SAVE_PATH = \"/kaggle/working/best_levit_model_v3.pth\"\nPLOT_SAVE_PATH = \"/kaggle/working/training_history_v3.png\"\nLOG_FILE_PATH = \"/kaggle/working/training_log_v3.csv\"\n\nBATCH_SIZE = 32\nEPOCHS_PHASE_1 = 10   \nEPOCHS_PHASE_2 = 90  \nTOTAL_EPOCHS = EPOCHS_PHASE_1 + EPOCHS_PHASE_2\nPATIENCE = 15         \n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Đang sử dụng thiết bị: {device}\")","metadata":{},"outputs":[],"execution_count":null},{"id":"68bf3b60","cell_type":"code","source":"# ==========================================\n# 2. CHUẨN BỊ DỮ LIỆU (DATASET)\n# ==========================================\nclass MalwareDataset(Dataset):\n    def __init__(self, dataframe, img_dir, transform=None):\n        self.dataframe = dataframe\n        self.img_dir = img_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, idx):\n        img_name = str(self.dataframe.iloc[idx, 0]) + \".png\"\n        img_path = os.path.join(self.img_dir, img_name)\n        \n        image = Image.open(img_path).convert('RGB')\n        label = int(self.dataframe.iloc[idx, 1]) - 1 \n\n        if self.transform:\n            image = self.transform(image)\n\n        return image, label","metadata":{},"outputs":[],"execution_count":null},{"id":"341ac5a1","cell_type":"code","source":"# 1. Chia tập dữ liệu\ndf = pd.read_csv(CSV_PATH)\ntrain_df, val_df = train_test_split(df, test_size=0.2, random_state=42, shuffle=True, stratify=df['Class'])\n\n# ==========================================\n# CẢI TIẾN: TÍNH MEAN VÀ STD CỦA TẬP BIG 2015\n# ==========================================\nprint(\"Đang tính toán Mean và Std cho tập dữ liệu BIG 2015...\")\n# Khởi tạo một Dataset tạm thời CHỈ dùng ToTensor (để chuyển pixel về dải 0-1)\ntemp_transform = transforms.Compose([\n    transforms.Resize((224, 224), interpolation=InterpolationMode.NEAREST),\n    transforms.ToTensor()\n])\ntemp_dataset = MalwareDataset(train_df, IMG_DIR, transform=temp_transform)\ntemp_loader = DataLoader(temp_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=2)\n\nchannels_sum, channels_squared_sum, num_batches = 0, 0, 0\nfor data, _ in tqdm(temp_loader, desc=\"Calc Mean/Std\", leave=False):\n    # Dữ liệu có shape: (batch_size, channels, height, width)\n    channels_sum += torch.mean(data, dim=[0, 2, 3])\n    channels_squared_sum += torch.mean(data**2, dim=[0, 2, 3])\n    num_batches += 1\n\n# Tính toán Mean và Std thực tế\nbig2015_mean = (channels_sum / num_batches).tolist()\nbig2015_std = ((channels_squared_sum / num_batches - torch.tensor(big2015_mean)**2)**0.5).tolist()\n\nprint(f\"-> Mean thực tế: {big2015_mean}\")\nprint(f\"-> Std thực tế: {big2015_std}\")\n\n# ==========================================\n# ÁP DỤNG TRANSFORM CHÍNH THỨC\n# ==========================================\ntrain_transform = transforms.Compose([\n    transforms.Resize((224, 224), interpolation=InterpolationMode.NEAREST),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=big2015_mean, std=big2015_std), # Sử dụng thông số vừa tính\n])\n\nval_transform = transforms.Compose([\n    transforms.Resize((224, 224), interpolation=InterpolationMode.NEAREST),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=big2015_mean, std=big2015_std), # Sử dụng thông số vừa tính\n])\n\n# Khởi tạo Dataset và DataLoader chính thức\ntrain_dataset = MalwareDataset(train_df, IMG_DIR, transform=train_transform)\nval_dataset = MalwareDataset(val_df, IMG_DIR, transform=val_transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=2)\nval_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=2)\nprint(\"Hoàn tất chuẩn bị DataLoader!\")","metadata":{},"outputs":[],"execution_count":null},{"id":"655cf8df","cell_type":"code","source":"# ==========================================\n# 3. KHỞI TẠO MÔ HÌNH & HÀM LOSS\n# ==========================================\nclass LeViT_Malware(nn.Module):\n    def __init__(self, model_name='levit_192', num_classes=9, pretrained=True):\n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=pretrained, num_classes=0)\n        num_features = self.backbone.num_features\n        \n        # --- CẢI TIẾN: Giảm cấu hình Head để tránh over-parameterized ---\n        self.head = nn.Sequential(\n            nn.BatchNorm1d(num_features),\n            nn.Dropout(p=0.2), \n            nn.Linear(num_features, 256), \n            nn.GELU(),\n            nn.BatchNorm1d(256),\n            nn.Dropout(p=0.1), \n            nn.Linear(256, num_classes)\n        )\n\n    def forward(self, x):\n        features = self.backbone(x)\n        return self.head(features)\n\nmodel = LeViT_Malware().to(device)\n\n# --- ROLLBACK: Bỏ Label Smoothing và Class Weights để Log Loss đạt tối đa ---\ncriterion = nn.CrossEntropyLoss()","metadata":{},"outputs":[],"execution_count":null},{"id":"f7b8327b","cell_type":"code","source":"# ==========================================\n# 4. EARLY STOPPING CLASS\n# ==========================================\nclass EarlyStopping:\n    def __init__(self, patience=5, path=MODEL_SAVE_PATH):\n        self.patience = patience\n        self.path = path\n        self.counter = 0\n        self.best_acc = 0.0\n        self.early_stop = False\n\n    def __call__(self, val_acc, model):\n        if val_acc > self.best_acc:\n            print(f\"Validation Accuracy tăng ({self.best_acc:.4f} --> {val_acc:.4f}). Lưu mô hình...\")\n            self.best_acc = val_acc\n            torch.save(model.state_dict(), self.path)\n            self.counter = 0\n        else:\n            self.counter += 1\n            print(f\"Early Stopping counter: {self.counter} out of {self.patience}\")\n            if self.counter >= self.patience:\n                self.early_stop = True\n\nearly_stopping = EarlyStopping(patience=PATIENCE)","metadata":{},"outputs":[],"execution_count":null},{"id":"1b78e3f3","cell_type":"code","source":"# ==========================================\n# 5. VÒNG LẶP HUẤN LUYỆN (TRAINING PIPELINE)\n# ==========================================\nhistory = {'train_loss': [], 'val_loss': [], 'train_acc': [], 'val_acc': []}\n\ndef train_epoch(model, dataloader, optimizer, criterion):\n    model.train()\n    running_loss, correct, total = 0.0, 0, 0\n    for images, labels in tqdm(dataloader, desc=\"Training\", leave=False):\n        images, labels = images.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        \n        # Giữ lại: Gradient Clipping \n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        \n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n        \n    return running_loss / len(dataloader), correct / total\n\ndef val_epoch(model, dataloader, criterion):\n    model.eval()\n    running_loss, correct, total = 0.0, 0, 0\n    with torch.no_grad():\n        for images, labels in tqdm(dataloader, desc=\"Validating\", leave=False):\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n            \n    return running_loss / len(dataloader), correct / total\n\n# Khởi tạo file log\nwith open(LOG_FILE_PATH, mode='w', newline='') as file:\n    writer = csv.writer(file)\n    writer.writerow(['Epoch', 'Train_Loss', 'Train_Acc', 'Val_Loss', 'Val_Acc'])\n\nprint(\"--- BẮT ĐẦU TRAINING ---\")\nfor epoch in range(TOTAL_EPOCHS):\n    if epoch == 0:\n        print(\"\\n[PHASE 1] Đóng băng Backbone, chỉ train Custom Head\")\n        for param in model.backbone.parameters():\n            param.requires_grad = False\n        optimizer = optim.AdamW(model.head.parameters(), lr=1e-3, weight_decay=0.05)\n        \n    elif epoch == EPOCHS_PHASE_1:\n        print(\"\\n[PHASE 2] Mở băng toàn bộ mạng (Fine-tuning)\")\n        for param in model.backbone.parameters():\n            param.requires_grad = True\n        optimizer = optim.AdamW([\n            {'params': model.backbone.parameters(), 'lr': 1e-5},\n            {'params': model.head.parameters(), 'lr': 1e-4}\n        ], weight_decay=1e-2)\n        \n        # Điều chỉnh eta_min lên 1e-6 cho đỡ sâu quá\n        scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS_PHASE_2, eta_min=1e-6)\n\n    print(f\"\\nEpoch {epoch+1}/{TOTAL_EPOCHS}\")\n    train_loss, train_acc = train_epoch(model, train_loader, optimizer, criterion)\n    val_loss, val_acc = val_epoch(model, val_loader, criterion)\n    \n    if epoch >= EPOCHS_PHASE_1:\n        scheduler.step()\n    \n    history['train_loss'].append(train_loss)\n    history['val_loss'].append(val_loss)\n    history['train_acc'].append(train_acc)\n    history['val_acc'].append(val_acc)\n    \n    print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f}\")\n    print(f\"Val Loss:   {val_loss:.4f} | Val Acc:   {val_acc:.4f}\")\n    \n    with open(LOG_FILE_PATH, mode='a', newline='') as file:\n        writer = csv.writer(file)\n        writer.writerow([epoch + 1, train_loss, train_acc, val_loss, val_acc])\n    \n    early_stopping(val_acc, model)\n    if early_stopping.early_stop:\n        print(\"Đã đạt giới hạn Early Stopping. Dừng huấn luyện sớm!\")\n        break","metadata":{},"outputs":[],"execution_count":null},{"id":"bfc8d51d","cell_type":"code","source":"# ==========================================\n# 6. VẼ VÀ LƯU ĐỒ THỊ\n# ==========================================\nepochs_run = len(history['train_loss'])\nepochs_range = range(1, epochs_run + 1)\n\nplt.figure(figsize=(14, 5))\n\nplt.subplot(1, 2, 1)\nplt.plot(epochs_range, history['train_acc'], label='Train Accuracy', marker='o')\nplt.plot(epochs_range, history['val_acc'], label='Validation Accuracy', marker='o')\nplt.title('Model Accuracy')\nplt.xlabel('Epochs')\nplt.ylabel('Accuracy')\nplt.legend()\nplt.grid(True)\n\nplt.subplot(1, 2, 2)\nplt.plot(epochs_range, history['train_loss'], label='Train Loss', marker='o')\nplt.plot(epochs_range, history['val_loss'], label='Validation Loss', marker='o')\nplt.title('Model Loss')\nplt.xlabel('Epochs')\nplt.ylabel('Loss')\nplt.legend()\nplt.grid(True)\n\nplt.tight_layout()\nplt.savefig(PLOT_SAVE_PATH) \nplt.show()\n\nprint(f\"\\nTraining hoàn tất! Trọng số tốt nhất lưu tại: {MODEL_SAVE_PATH}\")","metadata":{},"outputs":[],"execution_count":null}]}