{ "cells": [ { "cell_type": "code", "execution_count": 5, "metadata": { "id": "DDjtMfKwauQl" }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "Requirement already satisfied: timm in /venv/main/lib/python3.12/site-packages (1.0.25)\n", "Requirement already satisfied: gdown in /venv/main/lib/python3.12/site-packages (5.2.1)\n", "Collecting tensorboard\n", " Downloading tensorboard-2.20.0-py3-none-any.whl.metadata (1.8 kB)\n", "Requirement already satisfied: torch in /venv/main/lib/python3.12/site-packages (from timm) (2.10.0+cu130)\n", "Requirement already satisfied: torchvision in /venv/main/lib/python3.12/site-packages (from timm) (0.25.0+cu130)\n", "Requirement already satisfied: pyyaml in /venv/main/lib/python3.12/site-packages (from timm) (6.0.3)\n", "Requirement already satisfied: huggingface_hub in /venv/main/lib/python3.12/site-packages (from timm) (1.2.3)\n", "Requirement already satisfied: safetensors in /venv/main/lib/python3.12/site-packages (from timm) (0.7.0)\n", "Requirement already satisfied: beautifulsoup4 in /venv/main/lib/python3.12/site-packages (from gdown) (4.14.3)\n", "Requirement already satisfied: filelock in /venv/main/lib/python3.12/site-packages (from gdown) (3.20.1)\n", "Requirement already satisfied: requests[socks] in /venv/main/lib/python3.12/site-packages (from gdown) (2.32.5)\n", "Requirement already satisfied: tqdm in /venv/main/lib/python3.12/site-packages (from gdown) (4.67.1)\n", "Collecting absl-py>=0.4 (from tensorboard)\n", " Downloading absl_py-2.4.0-py3-none-any.whl.metadata (3.3 kB)\n", "Collecting grpcio>=1.48.2 (from tensorboard)\n", " Downloading grpcio-1.78.0-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.whl.metadata (3.8 kB)\n", "Collecting markdown>=2.6.8 (from tensorboard)\n", " Downloading markdown-3.10.2-py3-none-any.whl.metadata (5.1 kB)\n", "Requirement already satisfied: numpy>=1.12.0 in /venv/main/lib/python3.12/site-packages (from tensorboard) (2.4.1)\n", "Requirement already satisfied: packaging in /venv/main/lib/python3.12/site-packages (from tensorboard) (25.0)\n", "Requirement already satisfied: pillow in /venv/main/lib/python3.12/site-packages (from tensorboard) (12.1.0)\n", "Collecting protobuf!=4.24.0,>=3.19.6 (from tensorboard)\n", " Downloading protobuf-7.34.0-cp310-abi3-manylinux2014_x86_64.whl.metadata (595 bytes)\n", "Requirement already satisfied: setuptools>=41.0.0 in /venv/main/lib/python3.12/site-packages (from tensorboard) (80.9.0)\n", "Collecting tensorboard-data-server<0.8.0,>=0.7.0 (from tensorboard)\n", " Downloading tensorboard_data_server-0.7.2-py3-none-manylinux_2_31_x86_64.whl.metadata (1.1 kB)\n", "Collecting werkzeug>=1.0.1 (from tensorboard)\n", " Downloading werkzeug-3.1.6-py3-none-any.whl.metadata (4.0 kB)\n", "Requirement already satisfied: typing-extensions~=4.12 in /venv/main/lib/python3.12/site-packages (from grpcio>=1.48.2->tensorboard) (4.15.0)\n", "Requirement already satisfied: markupsafe>=2.1.1 in /venv/main/lib/python3.12/site-packages (from werkzeug>=1.0.1->tensorboard) (3.0.3)\n", "Requirement already satisfied: soupsieve>=1.6.1 in /venv/main/lib/python3.12/site-packages (from beautifulsoup4->gdown) (2.8.3)\n", "Requirement already satisfied: fsspec>=2023.5.0 in /venv/main/lib/python3.12/site-packages (from huggingface_hub->timm) (2025.12.0)\n", "Requirement already satisfied: hf-xet<2.0.0,>=1.2.0 in /venv/main/lib/python3.12/site-packages (from huggingface_hub->timm) (1.2.0)\n", "Requirement already satisfied: httpx<1,>=0.23.0 in /venv/main/lib/python3.12/site-packages (from huggingface_hub->timm) (0.28.1)\n", "Requirement already satisfied: shellingham in /venv/main/lib/python3.12/site-packages (from huggingface_hub->timm) (1.5.4)\n", "Requirement already satisfied: typer-slim in /venv/main/lib/python3.12/site-packages (from huggingface_hub->timm) (0.21.0)\n", "Requirement already satisfied: anyio in /venv/main/lib/python3.12/site-packages (from httpx<1,>=0.23.0->huggingface_hub->timm) (4.12.0)\n", "Requirement already satisfied: certifi in /venv/main/lib/python3.12/site-packages (from httpx<1,>=0.23.0->huggingface_hub->timm) (2025.11.12)\n", "Requirement already satisfied: httpcore==1.* in /venv/main/lib/python3.12/site-packages (from httpx<1,>=0.23.0->huggingface_hub->timm) (1.0.9)\n", "Requirement already satisfied: idna in /venv/main/lib/python3.12/site-packages (from httpx<1,>=0.23.0->huggingface_hub->timm) (3.11)\n", "Requirement already satisfied: h11>=0.16 in /venv/main/lib/python3.12/site-packages (from httpcore==1.*->httpx<1,>=0.23.0->huggingface_hub->timm) (0.16.0)\n", "Requirement already satisfied: charset_normalizer<4,>=2 in /venv/main/lib/python3.12/site-packages (from requests[socks]->gdown) (3.4.4)\n", "Requirement already satisfied: urllib3<3,>=1.21.1 in /venv/main/lib/python3.12/site-packages (from requests[socks]->gdown) (2.6.3)\n", "Requirement already satisfied: PySocks!=1.5.7,>=1.5.6 in /venv/main/lib/python3.12/site-packages (from requests[socks]->gdown) (1.7.1)\n", "Requirement already satisfied: sympy>=1.13.3 in /venv/main/lib/python3.12/site-packages (from torch->timm) (1.14.0)\n", "Requirement already satisfied: networkx>=2.5.1 in /venv/main/lib/python3.12/site-packages (from torch->timm) (3.6.1)\n", "Requirement already satisfied: jinja2 in /venv/main/lib/python3.12/site-packages (from torch->timm) (3.1.6)\n", "Requirement already satisfied: cuda-bindings==13.0.3 in /venv/main/lib/python3.12/site-packages (from torch->timm) (13.0.3)\n", "Requirement already satisfied: nvidia-cuda-nvrtc==13.0.88 in /venv/main/lib/python3.12/site-packages (from torch->timm) (13.0.88)\n", "Requirement already satisfied: nvidia-cuda-runtime==13.0.96 in /venv/main/lib/python3.12/site-packages (from torch->timm) (13.0.96)\n", "Requirement already satisfied: nvidia-cuda-cupti==13.0.85 in /venv/main/lib/python3.12/site-packages (from torch->timm) (13.0.85)\n", "Requirement already satisfied: nvidia-cudnn-cu13==9.15.1.9 in /venv/main/lib/python3.12/site-packages (from torch->timm) (9.15.1.9)\n", "Requirement already satisfied: nvidia-cublas==13.1.0.3 in /venv/main/lib/python3.12/site-packages (from torch->timm) (13.1.0.3)\n", "Requirement already satisfied: nvidia-cufft==12.0.0.61 in /venv/main/lib/python3.12/site-packages (from torch->timm) (12.0.0.61)\n", "Requirement already satisfied: nvidia-curand==10.4.0.35 in /venv/main/lib/python3.12/site-packages (from torch->timm) (10.4.0.35)\n", "Requirement already satisfied: nvidia-cusolver==12.0.4.66 in /venv/main/lib/python3.12/site-packages (from torch->timm) (12.0.4.66)\n", "Requirement already satisfied: nvidia-cusparse==12.6.3.3 in /venv/main/lib/python3.12/site-packages (from torch->timm) (12.6.3.3)\n", "Requirement already satisfied: nvidia-cusparselt-cu13==0.8.0 in /venv/main/lib/python3.12/site-packages (from torch->timm) (0.8.0)\n", "Requirement already satisfied: nvidia-nccl-cu13==2.28.9 in /venv/main/lib/python3.12/site-packages (from torch->timm) (2.28.9)\n", "Requirement already satisfied: nvidia-nvshmem-cu13==3.4.5 in /venv/main/lib/python3.12/site-packages (from torch->timm) (3.4.5)\n", "Requirement already satisfied: nvidia-nvtx==13.0.85 in /venv/main/lib/python3.12/site-packages (from torch->timm) (13.0.85)\n", "Requirement already satisfied: nvidia-nvjitlink==13.0.88 in /venv/main/lib/python3.12/site-packages (from torch->timm) (13.0.88)\n", "Requirement already satisfied: nvidia-cufile==1.15.1.6 in /venv/main/lib/python3.12/site-packages (from torch->timm) (1.15.1.6)\n", "Requirement already satisfied: triton==3.6.0 in /venv/main/lib/python3.12/site-packages (from torch->timm) (3.6.0)\n", "Requirement already satisfied: cuda-pathfinder~=1.1 in /venv/main/lib/python3.12/site-packages (from cuda-bindings==13.0.3->torch->timm) (1.3.3)\n", "Requirement already satisfied: mpmath<1.4,>=1.1.0 in /venv/main/lib/python3.12/site-packages (from sympy>=1.13.3->torch->timm) (1.3.0)\n", "Requirement already satisfied: click>=8.0.0 in /venv/main/lib/python3.12/site-packages (from typer-slim->huggingface_hub->timm) (8.3.1)\n", "Downloading tensorboard-2.20.0-py3-none-any.whl (5.5 MB)\n", "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m5.5/5.5 MB\u001b[0m \u001b[31m12.0 MB/s\u001b[0m \u001b[33m0:00:00\u001b[0mm0:00:01\u001b[0m00:01\u001b[0m\n", "\u001b[?25hDownloading tensorboard_data_server-0.7.2-py3-none-manylinux_2_31_x86_64.whl (6.6 MB)\n", "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m6.6/6.6 MB\u001b[0m \u001b[31m20.7 MB/s\u001b[0m \u001b[33m0:00:00\u001b[0mm0:00:01\u001b[0m00:01\u001b[0m\n", "\u001b[?25hDownloading absl_py-2.4.0-py3-none-any.whl (135 kB)\n", "Downloading grpcio-1.78.0-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.whl (6.7 MB)\n", "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m6.7/6.7 MB\u001b[0m \u001b[31m20.8 MB/s\u001b[0m \u001b[33m0:00:00\u001b[0mm0:00:01\u001b[0m00:01\u001b[0m\n", "\u001b[?25hDownloading markdown-3.10.2-py3-none-any.whl (108 kB)\n", "Downloading protobuf-7.34.0-cp310-abi3-manylinux2014_x86_64.whl (324 kB)\n", "Downloading werkzeug-3.1.6-py3-none-any.whl (225 kB)\n", "Installing collected packages: werkzeug, tensorboard-data-server, protobuf, markdown, grpcio, absl-py, tensorboard\n", "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m7/7\u001b[0m [tensorboard]\u001b[0m [tensorboard]\n", "\u001b[1A\u001b[2KSuccessfully installed absl-py-2.4.0 grpcio-1.78.0 markdown-3.10.2 protobuf-7.34.0 tensorboard-2.20.0 tensorboard-data-server-0.7.2 werkzeug-3.1.6\n", "\u001b[33mWARNING: Running pip as the 'root' user can result in broken permissions and conflicting behaviour with the system package manager, possibly rendering your system unusable. It is recommended to use a virtual environment instead: https://pip.pypa.io/warnings/venv. Use the --root-user-action option if you know what you are doing and want to suppress this warning.\u001b[0m\u001b[33m\n", "\u001b[0mNote: you may need to restart the kernel to use updated packages.\n" ] } ], "source": [ "%pip install timm gdown tensorboard" ] }, { "cell_type": "code", "execution_count": 1, "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "Package Version\n", "----------------------- ------------\n", "anyio 4.12.0\n", "asttokens 3.0.1\n", "certifi 2025.11.12\n", "charset-normalizer 3.4.4\n", "click 8.3.1\n", "comm 0.2.3\n", "cuda-bindings 13.0.3\n", "cuda-pathfinder 1.3.3\n", "debugpy 1.8.19\n", "decorator 5.2.1\n", "executing 2.2.1\n", "filelock 3.20.1\n", "fsspec 2025.12.0\n", "h11 0.16.0\n", "hf-xet 1.2.0\n", "httpcore 1.0.9\n", "httpx 0.28.1\n", "huggingface_hub 1.2.3\n", "idna 3.11\n", "ipykernel 7.1.0\n", "ipython 9.8.0\n", "ipython_pygments_lexers 1.1.1\n", "ipywidgets 8.1.8\n", "jedi 0.19.2\n", "Jinja2 3.1.6\n", "jupyter_client 8.7.0\n", "jupyter_core 5.9.1\n", "jupyterlab_widgets 3.0.16\n", "MarkupSafe 3.0.3\n", "matplotlib-inline 0.2.1\n", "mpmath 1.3.0\n", "nest-asyncio 1.6.0\n", "networkx 3.6.1\n", "numpy 2.4.1\n", "nvidia-cublas 13.1.0.3\n", "nvidia-cuda-cupti 13.0.85\n", "nvidia-cuda-nvrtc 13.0.88\n", "nvidia-cuda-runtime 13.0.96\n", "nvidia-cudnn-cu13 9.15.1.9\n", "nvidia-cufft 12.0.0.61\n", "nvidia-cufile 1.15.1.6\n", "nvidia-curand 10.4.0.35\n", "nvidia-cusolver 12.0.4.66\n", "nvidia-cusparse 12.6.3.3\n", "nvidia-cusparselt-cu13 0.8.0\n", "nvidia-nccl-cu13 2.28.9\n", "nvidia-nvjitlink 13.0.88\n", "nvidia-nvshmem-cu13 3.4.5\n", "nvidia-nvtx 13.0.85\n", "packaging 25.0\n", "parso 0.8.5\n", "pexpect 4.9.0\n", "pillow 12.1.0\n", "pip 25.3\n", "platformdirs 4.5.1\n", "prompt_toolkit 3.0.52\n", "psutil 7.2.1\n", "ptyprocess 0.7.0\n", "pure_eval 0.2.3\n", "Pygments 2.19.2\n", "python-dateutil 2.9.0.post0\n", "PyYAML 6.0.3\n", "pyzmq 27.1.0\n", "requests 2.32.5\n", "sentencepiece 0.2.1\n", "setuptools 80.9.0\n", "shellingham 1.5.4\n", "six 1.17.0\n", "stack-data 0.6.3\n", "sympy 1.14.0\n", "torch 2.10.0+cu130\n", "torchaudio 2.10.0+cu130\n", "torchcodec 0.10.0\n", "torchdata 0.10.0\n", "torchtext 0.6.0\n", "torchvision 0.25.0+cu130\n", "tornado 6.5.4\n", "tqdm 4.67.1\n", "traitlets 5.14.3\n", "triton 3.6.0\n", "typer-slim 0.21.0\n", "typing_extensions 4.15.0\n", "urllib3 2.6.3\n", "wcwidth 0.2.14\n", "wheel 0.45.1\n", "widgetsnbextension 4.0.15\n", "Note: you may need to restart the kernel to use updated packages.\n" ] } ], "source": [ "%pip list" ] }, { "cell_type": "code", "execution_count": 6, "metadata": { "id": "xJYUsKdBCPVS" }, "outputs": [], "source": [ "import os\n", "import shutil\n", "import torch\n", "import torch.nn as nn\n", "import torch.optim as optim\n", "from torchvision import datasets, transforms\n", "from timm import create_model\n", "from torch.optim.lr_scheduler import CosineAnnealingLR\n", "from torch.utils.data import DataLoader\n", "from torch.utils.tensorboard import SummaryWriter\n", "from tqdm import tqdm # For progress bar\n", "from torchvision.transforms import RandAugment\n", "from timm.data import Mixup\n", "from timm.loss import SoftTargetCrossEntropy\n", "from timm.layers import DropPath # Updated import path\n", "from timm.scheduler.cosine_lr import CosineLRScheduler\n", "\n" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "import gdown\n", "\n", "url = 'https://drive.google.com/'\n", "output = 'dat2.zip'\n", "gdown.download(url, output, quiet=False, fuzzy=True)" ] }, { "cell_type": "code", "execution_count": 10, "metadata": { "colab": { "base_uri": "https://localhost:8080/" }, "id": "cbW5kekcM0nQ", "outputId": "0c0d198c-d492-4b04-8870-eea889d8dfea" }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "Files extracted to: /workspace/red\n" ] } ], "source": [ "import zipfile\n", "import os\n", "\n", "zip_file_name = '/workspace/dat2.zip'\n", "extract_dir = '/workspace/red' # Target directory for extraction\n", "\n", "# Create the target directory if it doesn't exist\n", "if not os.path.exists(extract_dir):\n", " os.makedirs(extract_dir)\n", "\n", "with zipfile.ZipFile(zip_file_name, 'r') as zip_ref:\n", " # Extract all contents to the specified directory\n", " zip_ref.extractall(extract_dir)\n", "\n", "print(f\"Files extracted to: {extract_dir}\")" ] }, { "cell_type": "code", "execution_count": 11, "metadata": { "id": "5D1NQo1LCStD" }, "outputs": [], "source": [ "# Paths and Constants\n", "data_dir = \"/workspace/red/splits1\"\n", "num_classes = 7\n", "batch_size = 128 # Adjust based on GPU memory\n", "num_epochs = 50 # Increased number of epochs for better convergence\n", "learning_rate = 5e-4 # Lowered learning rate for fine-tuning\n", "weight_decay = 0.01 # Adjusted weight decay\n", "image_size = 224\n", "log_interval = 100 # Log metrics every 100 batches\n" ] }, { "cell_type": "code", "execution_count": 12, "metadata": { "id": "JgcWhUooCfC7" }, "outputs": [], "source": [ "transform_train = transforms.Compose([\n", " transforms.Resize((image_size, image_size), interpolation=transforms.InterpolationMode.BICUBIC),\n", " RandAugment(), # Enhanced augmentation\n", " transforms.ToTensor(),\n", " transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),\n", " transforms.RandomErasing(p=0.1),\n", "])\n" ] }, { "cell_type": "code", "execution_count": 13, "metadata": { "id": "5YWSDRJrCgwF" }, "outputs": [], "source": [ "transform_test = transforms.Compose([\n", " transforms.Resize((image_size, image_size), interpolation=transforms.InterpolationMode.BICUBIC),\n", " transforms.ToTensor(),\n", " transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),\n", "])\n" ] }, { "cell_type": "code", "execution_count": 14, "metadata": { "id": "sqcG-hMtChPa" }, "outputs": [], "source": [ "train_dataset = datasets.ImageFolder(os.path.join(data_dir, 'train'), transform=transform_train)\n", "train_loader = DataLoader(\n", " train_dataset,\n", " batch_size=batch_size,\n", " shuffle=True,\n", " num_workers=16, # Reduced from 8 to 2\n", " pin_memory=True,\n", " prefetch_factor=4,\n", " persistent_workers=True\n", ")\n" ] }, { "cell_type": "code", "execution_count": 15, "metadata": { "id": "h2QXKLTqCmT8" }, "outputs": [], "source": [ "val_images_dir = os.path.join(\"/workspace/red/splits1/val\", '')\n", "val_dataset = datasets.ImageFolder(val_images_dir, transform=transform_test)\n", "val_loader = DataLoader(\n", " val_dataset,\n", " batch_size=batch_size,\n", " shuffle=False,\n", " num_workers=16, # Reduced from 8 to 2\n", " pin_memory=True,\n", " prefetch_factor=4,\n", " persistent_workers=True\n", ")\n" ] }, { "cell_type": "code", "execution_count": 16, "metadata": { "colab": { "base_uri": "https://localhost:8080/" }, "id": "EQzvAErvCpJM", "outputId": "704d863e-3bdf-487f-e62e-8c5f724ed99f" }, "outputs": [ { "data": { "application/vnd.jupyter.widget-view+json": { "model_id": "50ffe1bd5e6443719a4752151afd6b14", "version_major": 2, "version_minor": 0 }, "text/plain": [ "model.safetensors: 0%| | 0.00/88.2M [00:00 1:\n", " print(f\"Using {torch.cuda.device_count()} GPUs\")\n", " model = nn.DataParallel(model) # This will use all available GPUs" ] }, { "cell_type": "code", "execution_count": 22, "metadata": { "id": "W-stw1OBDLqp" }, "outputs": [], "source": [ "model = model.to(device)\n" ] }, { "cell_type": "code", "execution_count": 23, "metadata": {}, "outputs": [], "source": [ "class_weights = torch.tensor([1.560, 3.737, 2.242, 0.541, 0.527, 0.970, 1.149], dtype=torch.float)\n", "class_weights = class_weights.to(device)" ] }, { "cell_type": "code", "execution_count": 24, "metadata": { "id": "tI3m3KquDPXI" }, "outputs": [], "source": [ "criterion = nn.CrossEntropyLoss(weight=class_weights) # For Mixup and CutMix\n" ] }, { "cell_type": "code", "execution_count": 25, "metadata": { "id": "afX6EEL8DQvq" }, "outputs": [], "source": [ "optimizer = optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=weight_decay)" ] }, { "cell_type": "code", "execution_count": 26, "metadata": { "id": "Yf0JflY9DR_t" }, "outputs": [], "source": [ "scheduler = CosineLRScheduler(\n", " optimizer,\n", " t_initial=num_epochs,\n", " lr_min=1e-5,\n", " warmup_t=5,\n", " warmup_lr_init=1e-6,\n", " warmup_prefix=True \n", ")" ] }, { "cell_type": "code", "execution_count": 27, "metadata": { "id": "jmO4QHy3DT0j" }, "outputs": [], "source": [ "scaler = torch.amp.GradScaler(device='cuda') # Updated instantiation\n" ] }, { "cell_type": "code", "execution_count": 28, "metadata": { "id": "pW1uIqzeDVZv" }, "outputs": [], "source": [ "writer = SummaryWriter() # For TensorBoard logging\n" ] }, { "cell_type": "code", "execution_count": 29, "metadata": { "id": "eFQv3QNYDWyq" }, "outputs": [], "source": [ "def train_one_epoch(epoch):\n", " model.train()\n", " running_loss, correct, total = 0.0, 0, 0\n", "\n", " # Progress bar for training loop\n", " train_loader_tqdm = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs} [Training]\", leave=False)\n", " for batch_idx, (images, labels) in enumerate(train_loader_tqdm):\n", " images, labels = images.to(device, non_blocking=True), labels.to(device, non_blocking=True)\n", "\n", " optimizer.zero_grad()\n", "\n", " with torch.cuda.amp.autocast():\n", " outputs = model(images)\n", " loss = criterion(outputs, labels)\n", "\n", " scaler.scale(loss).backward()\n", " scaler.step(optimizer)\n", " scaler.update()\n", "\n", " running_loss += loss.item() * images.size(0)\n", " total += labels.size(0)\n", "\n", " # Since labels are soft, calculate accuracy based on predicted class vs hard labels\n", " _, predicted = outputs.max(1)\n", " correct += predicted.eq(labels).sum().item()\n", "\n", " # Update progress bar (accuracy in percentage)\n", " if (batch_idx + 1) % log_interval == 0 or (batch_idx + 1) == len(train_loader):\n", " current_loss = loss.item()\n", " current_acc = 100. * correct / total\n", " train_loader_tqdm.set_postfix(loss=f\"{current_loss:.4f}\", accuracy=f\"{current_acc:.2f}%\")\n", "\n", " epoch_loss = running_loss / total\n", " epoch_acc = 100. * correct / total # Multiply by 100 to get percentage\n", " writer.add_scalar('Loss/train', epoch_loss, epoch)\n", " writer.add_scalar('Accuracy/train', epoch_acc, epoch)\n", " print(f\"Epoch [{epoch+1}/{num_epochs}], Loss: {epoch_loss:.4f}, Acc: {epoch_acc:.2f}%\") # Acc in %\n" ] }, { "cell_type": "code", "execution_count": 30, "metadata": { "id": "Bi9RQ6CDDZuE" }, "outputs": [], "source": [ "def validate(epoch):\n", " model.eval()\n", " val_loss, correct, total = 0.0, 0, 0\n", "\n", " # Progress bar for validation loop\n", " val_loader_tqdm = tqdm(val_loader, desc=f\"Epoch {epoch+1}/{num_epochs} [Validation]\", leave=False)\n", "\n", " criterion_val = nn.CrossEntropyLoss() # Standard loss for validation\n", "\n", " with torch.no_grad():\n", " for batch_idx, (images, labels) in enumerate(val_loader_tqdm):\n", " images, labels = images.to(device, non_blocking=True), labels.to(device, non_blocking=True)\n", "\n", " with torch.cuda.amp.autocast():\n", " outputs = model(images)\n", " loss = criterion_val(outputs, labels)\n", "\n", " val_loss += loss.item() * images.size(0)\n", " total += labels.size(0)\n", " _, predicted = outputs.max(1)\n", " correct += predicted.eq(labels).sum().item()\n", "\n", " # Update progress bar (accuracy in percentage)\n", " if (batch_idx + 1) % log_interval == 0 or (batch_idx + 1) == len(val_loader):\n", " current_loss = loss.item()\n", " current_acc = 100. * correct / total\n", " val_loader_tqdm.set_postfix(loss=f\"{current_loss:.4f}\", accuracy=f\"{current_acc:.2f}%\")\n", "\n", " epoch_loss = val_loss / total\n", " epoch_acc = 100. * correct / total # Multiply by 100 to get percentage\n", " writer.add_scalar('Loss/val', epoch_loss, epoch)\n", " writer.add_scalar('Accuracy/val', epoch_acc, epoch)\n", " print(f\"Validation Loss: {epoch_loss:.4f}, Acc: {epoch_acc:.2f}%\") # Acc in %\n", "\n", " return epoch_acc" ] }, { "cell_type": "code", "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/", "height": 474 }, "id": "YJ53kRT3DcuJ", "outputId": "ef6e52b1-8f2a-43c3-a5f4-49009ce167d8" }, "outputs": [ { "name": "stderr", "output_type": "stream", "text": [ "Epoch 1/50 [Training]: 0%| | 0/420 [00:00 best_acc:\n", " best_acc = val_acc\n", " os.makedirs('./models', exist_ok=True)\n", " # If using DataParallel, save the underlying model\n", " if isinstance(model, nn.DataParallel):\n", " torch.save(model.module.state_dict(), './models/best_vit_tiny_imagenet.pth')\n", " else:\n", " torch.save(model.state_dict(), './models/best_vit_tiny_imagenet.pth')\n", " print(f\"New best model saved with accuracy: {best_acc:.2f}%\")\n", "\n", "print(\"Training complete. Best validation accuracy:\", best_acc)" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [] } ], "metadata": { "accelerator": "GPU", "colab": { "gpuType": "T4", "provenance": [] }, "kernelspec": { "display_name": "Python3 (main venv)", "language": "python", "name": "main" }, "language_info": { "codemirror_mode": { "name": "ipython", "version": 3 }, "file_extension": ".py", "mimetype": "text/x-python", "name": "python", "nbconvert_exporter": "python", "pygments_lexer": "ipython3", "version": "3.12.12" } }, "nbformat": 4, "nbformat_minor": 4 }