{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"}],"dockerImageVersionId":31192,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:56:23.822311Z","iopub.execute_input":"2025-12-05T12:56:23.822539Z","iopub.status.idle":"2025-12-05T12:56:29.716674Z","shell.execute_reply.started":"2025-12-05T12:56:23.822522Z","shell.execute_reply":"2025-12-05T12:56:29.715981Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:56:29.71805Z","iopub.execute_input":"2025-12-05T12:56:29.718385Z","iopub.status.idle":"2025-12-05T12:56:30.039267Z","shell.execute_reply.started":"2025-12-05T12:56:29.718359Z","shell.execute_reply":"2025-12-05T12:56:30.038526Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nroot = \"/kaggle/input\"\nfor path, dirs, files in os.walk(root):\n    print(path)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:56:30.040087Z","iopub.execute_input":"2025-12-05T12:56:30.040519Z","iopub.status.idle":"2025-12-05T12:56:30.057089Z","shell.execute_reply.started":"2025-12-05T12:56:30.040498Z","shell.execute_reply":"2025-12-05T12:56:30.056548Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nfor dirname, _, filenames in os.walk('/kaggle/input/UBC-OCEAN'):\n    print(dirname)\n    for filename in filenames:\n        print(\"  \", filename)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:56:30.057751Z","iopub.execute_input":"2025-12-05T12:56:30.057953Z","iopub.status.idle":"2025-12-05T12:56:30.088324Z","shell.execute_reply.started":"2025-12-05T12:56:30.057936Z","shell.execute_reply":"2025-12-05T12:56:30.087619Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls /kaggle/input","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:56:30.089212Z","iopub.execute_input":"2025-12-05T12:56:30.089461Z","iopub.status.idle":"2025-12-05T12:56:30.209416Z","shell.execute_reply.started":"2025-12-05T12:56:30.089431Z","shell.execute_reply":"2025-12-05T12:56:30.208678Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls /kaggle/input/UBC-OCEAN","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:56:30.210606Z","iopub.execute_input":"2025-12-05T12:56:30.211198Z","iopub.status.idle":"2025-12-05T12:56:30.33136Z","shell.execute_reply.started":"2025-12-05T12:56:30.211164Z","shell.execute_reply":"2025-12-05T12:56:30.330389Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\nbase = \"/kaggle/input/UBC-OCEAN\"\n\ntrain_df = pd.read_csv(f\"{base}/train.csv\")\ntest_df  = pd.read_csv(f\"{base}/test.csv\")\nsample_df = pd.read_csv(f\"{base}/sample_submission.csv\")\n\ntrain_df.head(), test_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:56:30.334139Z","iopub.execute_input":"2025-12-05T12:56:30.334382Z","iopub.status.idle":"2025-12-05T12:56:30.387496Z","shell.execute_reply.started":"2025-12-05T12:56:30.33435Z","shell.execute_reply":"2025-12-05T12:56:30.386746Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# Correct paths\ntrain_df = pd.read_csv('/kaggle/input/UBC-OCEAN/train.csv')\ntest_df  = pd.read_csv('/kaggle/input/UBC-OCEAN/test.csv')\n\ntrain_df.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:56:30.388368Z","iopub.execute_input":"2025-12-05T12:56:30.389221Z","iopub.status.idle":"2025-12-05T12:56:30.407286Z","shell.execute_reply.started":"2025-12-05T12:56:30.389194Z","shell.execute_reply":"2025-12-05T12:56:30.406766Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nbase = \"/kaggle/input/UBC-OCEAN/train_images\"\n\nfor root, dirs, files in os.walk(base):\n    print(root, \"->\", len(files), \"files\")\n    break\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:56:30.407983Z","iopub.execute_input":"2025-12-05T12:56:30.408227Z","iopub.status.idle":"2025-12-05T12:56:30.493879Z","shell.execute_reply.started":"2025-12-05T12:56:30.408201Z","shell.execute_reply":"2025-12-05T12:56:30.493223Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls /kaggle/input/UBC-OCEAN/train_images | head\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:56:30.494654Z","iopub.execute_input":"2025-12-05T12:56:30.495193Z","iopub.status.idle":"2025-12-05T12:56:30.619754Z","shell.execute_reply.started":"2025-12-05T12:56:30.495166Z","shell.execute_reply":"2025-12-05T12:56:30.618385Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_id = train_df.iloc[0][\"image_id\"]\npath = \"/kaggle/input/UBC-OCEAN/train_images\"\n\nprint(\"Searching for:\", image_id)\n\nmatches = []\nfor root, dirs, files in os.walk(path):\n    for f in files:\n        if f.startswith(str(image_id)):\n            matches.append(os.path.join(root, f))\n\nmatches[:10], len(matches)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:56:30.620883Z","iopub.execute_input":"2025-12-05T12:56:30.621187Z","iopub.status.idle":"2025-12-05T12:56:30.64413Z","shell.execute_reply.started":"2025-12-05T12:56:30.621155Z","shell.execute_reply":"2025-12-05T12:56:30.643428Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# Correct paths for the CSVs\ntrain_df = pd.read_csv('/kaggle/input/UBC-OCEAN/train.csv')\ntest_df = pd.read_csv('/kaggle/input/UBC-OCEAN/test.csv')\n\ntrain_df.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:56:30.644914Z","iopub.execute_input":"2025-12-05T12:56:30.645167Z","iopub.status.idle":"2025-12-05T12:56:30.660111Z","shell.execute_reply.started":"2025-12-05T12:56:30.645145Z","shell.execute_reply":"2025-12-05T12:56:30.659322Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cv2\nimport os\n\n# Get an example image_id from the train CSV\nimage_id = train_df.iloc[0][\"image_id\"]  # e.g., 4, 4608, etc.\n\n# Build correct path\nimage_path = f\"/kaggle/input/UBC-OCEAN/train_images/{image_id}.png\"\n\n# Load the image\nimg = cv2.imread(image_path)\n\nprint(\"Loaded:\", image_path)\nprint(\"Shape:\", img.shape)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:56:30.660801Z","iopub.execute_input":"2025-12-05T12:56:30.661035Z","iopub.status.idle":"2025-12-05T12:56:49.380599Z","shell.execute_reply.started":"2025-12-05T12:56:30.661014Z","shell.execute_reply":"2025-12-05T12:56:49.37977Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 1.1 Explore directory structure\nimport os\nfrom pprint import pprint\n\nbase = \"/kaggle/input/UBC-OCEAN\"\n\nprint(\"Base exists:\", os.path.exists(base))\nprint(\"\\nTop-level files/folders under /kaggle/input:\")\npprint(sorted(os.listdir(\"/kaggle/input\")))\n\nprint(\"\\nContents of UBC-OCEAN:\")\nfor root, dirs, files in os.walk(base):\n    # print first few directories and files only for readability\n    rel = os.path.relpath(root, base)\n    if rel == \".\":\n        rel = \"/\"\n    print(f\"\\nFolder: {rel}\")\n    if dirs:\n        print(\"  subdirs:\", dirs[:10])\n    if files:\n        print(\"  files (first 20):\", files[:20])\n    # stop descending after 3 levels for brevity\n    if root.count(os.sep) - base.count(os.sep) >= 2:\n        continue\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:56:49.381505Z","iopub.execute_input":"2025-12-05T12:56:49.381786Z","iopub.status.idle":"2025-12-05T12:56:50.108764Z","shell.execute_reply.started":"2025-12-05T12:56:49.381758Z","shell.execute_reply":"2025-12-05T12:56:50.10811Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 1.2 Load train/test CSVs and show quick EDA\nimport pandas as pd\n\nbase = \"/kaggle/input/UBC-OCEAN\"\ntrain_csv = os.path.join(base, \"train.csv\")\ntest_csv  = os.path.join(base, \"test.csv\")\nsample_csv = os.path.join(base, \"sample_submission.csv\")\n\nprint(\"train.csv exists:\", os.path.exists(train_csv))\nprint(\"test.csv exists: \", os.path.exists(test_csv))\nprint(\"sample_submission exists:\", os.path.exists(sample_csv))\n\ntrain_df = pd.read_csv(train_csv)\ntest_df  = pd.read_csv(test_csv)\nsample_df = pd.read_csv(sample_csv)\n\nprint(\"\\ntrain_df.shape:\", train_df.shape)\nprint(\"test_df.shape:\", test_df.shape)\nprint(\"\\ntrain_df.columns:\", train_df.columns.tolist())\nprint(\"\\nFirst 5 rows of train_df:\")\ndisplay(train_df.head())\n\n# If there is a class/label column, show distribution. Try common names.\npossible_label_cols = [c for c in train_df.columns if c.lower() in (\"label\",\"target\",\"class\",\"subtype\")]\nprint(\"\\nDetected possible label columns:\", possible_label_cols)\n\nif possible_label_cols:\n    lab = possible_label_cols[0]\n    print(f\"\\nLabel distribution for '{lab}':\")\n    display(train_df[lab].value_counts().sort_index())\nelse:\n    # fallback: try to infer multiclass columns\n    print(\"\\nNo obvious label column detected. Show first few columns and sample values to inspect.\")\n    display(train_df.iloc[:, :6].head())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:56:50.109468Z","iopub.execute_input":"2025-12-05T12:56:50.109722Z","iopub.status.idle":"2025-12-05T12:56:50.137765Z","shell.execute_reply.started":"2025-12-05T12:56:50.109704Z","shell.execute_reply":"2025-12-05T12:56:50.137018Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"try:\n    import timm\n    print(\"timm version:\", timm.__version__)\nexcept ImportError as e:\n    print(\"timm is NOT available:\", e)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:56:50.138672Z","iopub.execute_input":"2025-12-05T12:56:50.138908Z","iopub.status.idle":"2025-12-05T12:57:00.17179Z","shell.execute_reply.started":"2025-12-05T12:56:50.138891Z","shell.execute_reply":"2025-12-05T12:57:00.171085Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\n\nfrom PIL import Image, ImageFile\nImage.MAX_IMAGE_PIXELS = None\nImageFile.LOAD_TRUNCATED_IMAGES = True\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import LabelEncoder\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\n\nimport torchvision.transforms as T\nimport timm\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:57:00.172512Z","iopub.execute_input":"2025-12-05T12:57:00.172812Z","iopub.status.idle":"2025-12-05T12:57:01.008192Z","shell.execute_reply.started":"2025-12-05T12:57:00.172791Z","shell.execute_reply":"2025-12-05T12:57:01.007502Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torchvision.models as models\nfrom torchvision.models import inception_v3, swin_t\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:57:01.009053Z","iopub.execute_input":"2025-12-05T12:57:01.009571Z","iopub.status.idle":"2025-12-05T12:57:01.012919Z","shell.execute_reply.started":"2025-12-05T12:57:01.00955Z","shell.execute_reply":"2025-12-05T12:57:01.012382Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BASE_PATH = \"/kaggle/input/UBC-OCEAN\"\nIMAGE_DIR_TRAIN = f\"{BASE_PATH}/train_images\"\nIMAGE_DIR_TEST  = f\"{BASE_PATH}/test_images\"\n\ntrain_df = pd.read_csv(f\"{BASE_PATH}/train.csv\")\ntest_df  = pd.read_csv(f\"{BASE_PATH}/test.csv\")\n\nprint(\"train_df shape:\", train_df.shape)\nprint(\"test_df shape:\", test_df.shape)\nprint(\"Columns in train_df:\", train_df.columns.tolist())\n\ntrain_df.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:57:01.013722Z","iopub.execute_input":"2025-12-05T12:57:01.014189Z","iopub.status.idle":"2025-12-05T12:57:01.045481Z","shell.execute_reply.started":"2025-12-05T12:57:01.014163Z","shell.execute_reply":"2025-12-05T12:57:01.044921Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_col = \"label\"   # e.g. \"label\", \"subtype\", etc.\n\nassert label_col in train_df.columns, f\"{label_col} not found in train_df columns: {train_df.columns.tolist()}\"\n\nle = LabelEncoder()\ntrain_df[\"label_enc\"] = le.fit_transform(train_df[label_col])\n\nnum_classes = train_df[\"label_enc\"].nunique()\nprint(\"Number of classes:\", num_classes)\ntrain_df[[label_col, \"label_enc\"]].head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:57:01.046113Z","iopub.execute_input":"2025-12-05T12:57:01.046385Z","iopub.status.idle":"2025-12-05T12:57:01.059102Z","shell.execute_reply.started":"2025-12-05T12:57:01.046361Z","shell.execute_reply":"2025-12-05T12:57:01.05839Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df_split, val_df_split = train_test_split(\n    train_df,\n    test_size=0.2,\n    stratify=train_df[\"label_enc\"],\n    random_state=42\n)\n\nprint(\"Train size:\", len(train_df_split))\nprint(\"Val size:\", len(val_df_split))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:57:01.059918Z","iopub.execute_input":"2025-12-05T12:57:01.060434Z","iopub.status.idle":"2025-12-05T12:57:01.077582Z","shell.execute_reply.started":"2025-12-05T12:57:01.060408Z","shell.execute_reply":"2025-12-05T12:57:01.076851Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IMG_SIZE = 224\n\ntrain_transform = T.Compose([\n    T.Resize((IMG_SIZE, IMG_SIZE)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(10),\n    T.ToTensor(),\n    T.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    ),\n])\n\nval_transform = T.Compose([\n    T.Resize((IMG_SIZE, IMG_SIZE)),\n    T.ToTensor(),\n    T.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    ),\n])\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:57:01.080605Z","iopub.execute_input":"2025-12-05T12:57:01.080846Z","iopub.status.idle":"2025-12-05T12:57:01.091641Z","shell.execute_reply.started":"2025-12-05T12:57:01.080829Z","shell.execute_reply":"2025-12-05T12:57:01.091051Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class OCEANDataset(Dataset):\n    def __init__(self, df, image_dir, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.image_dir = image_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image_id = row[\"image_id\"]   # make sure this column exists in train.csv\n\n        img_path = os.path.join(self.image_dir, f\"{image_id}.png\")\n        if not os.path.exists(img_path):\n            raise FileNotFoundError(f\"Image not found: {img_path}\")\n\n        image = Image.open(img_path).convert(\"RGB\")\n\n        if self.transform is not None:\n            image = self.transform(image)\n\n        label = int(row[\"label_enc\"])\n        return image, label\n\n\nclass OCEANTestDataset(Dataset):\n    def __init__(self, df, image_dir, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.image_dir = image_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image_id = row[\"image_id\"]\n\n        img_path = os.path.join(self.image_dir, f\"{image_id}.png\")\n        if not os.path.exists(img_path):\n            raise FileNotFoundError(f\"Test image not found: {img_path}\")\n\n        image = Image.open(img_path).convert(\"RGB\")\n\n        if self.transform is not None:\n            image = self.transform(image)\n\n        return image, image_id\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:57:01.092236Z","iopub.execute_input":"2025-12-05T12:57:01.092492Z","iopub.status.idle":"2025-12-05T12:57:01.109833Z","shell.execute_reply.started":"2025-12-05T12:57:01.092466Z","shell.execute_reply":"2025-12-05T12:57:01.109266Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BATCH_SIZE = 4  # smaller batch for debugging\n\ntrain_dataset = OCEANDataset(train_df_split, IMAGE_DIR_TRAIN, transform=train_transform)\nval_dataset   = OCEANDataset(val_df_split,   IMAGE_DIR_TRAIN, transform=val_transform)\ntest_dataset  = OCEANTestDataset(test_df,    IMAGE_DIR_TEST,  transform=val_transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True,  num_workers=0)\nval_loader   = DataLoader(val_dataset,   batch_size=BATCH_SIZE, shuffle=False, num_workers=0)\ntest_loader  = DataLoader(test_dataset,  batch_size=BATCH_SIZE, shuffle=False, num_workers=0)\n\nprint(\"DataLoaders created.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:57:01.11062Z","iopub.execute_input":"2025-12-05T12:57:01.11098Z","iopub.status.idle":"2025-12-05T12:57:01.129055Z","shell.execute_reply.started":"2025-12-05T12:57:01.110962Z","shell.execute_reply":"2025-12-05T12:57:01.128272Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\nclass InceptionSwinHybrid(nn.Module):\n    def __init__(self, num_classes):\n        super().__init__()\n\n        # ---- InceptionV3 branch (base CNN) ----\n        # No pretrained weights (no internet needed)\n        self.inception = inception_v3(weights=None, aux_logits=False)\n        incep_in_features = self.inception.fc.in_features  # last FC input dim\n\n        # Remove original classification head\n        self.inception.fc = nn.Identity()\n        # Project Inception features to 512-dim embedding\n        self.incep_proj = nn.Linear(incep_in_features, 512)\n\n        # ---- Swin Transformer branch ----\n        self.swin = swin_t(weights=None)\n        swin_in_features = self.swin.head.in_features\n\n        # Remove Swin classifier head\n        self.swin.head = nn.Identity()\n        # Project Swin features to 512-dim embedding\n        self.swin_proj = nn.Linear(swin_in_features, 512)\n\n        # ---- Fusion + final classifier ----\n        # Concatenate [Inception(512), Swin(512)] -> 1024\n        self.classifier = nn.Linear(512 * 2, num_classes)\n\n    def forward(self, x):\n        # Inception branch\n        incep_feat = self.inception(x)            # [B, incep_in_features]\n        incep_feat = self.incep_proj(incep_feat)  # [B, 512]\n\n        # Swin branch\n        swin_feat = self.swin(x)                  # [B, swin_in_features]\n        swin_feat = self.swin_proj(swin_feat)     # [B, 512]\n\n        # Fuse\n        fused = torch.cat([incep_feat, swin_feat], dim=1)  # [B, 1024]\n\n        # Class logits\n        out = self.classifier(fused)              # [B, num_classes]\n        return out\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:57:01.130415Z","iopub.execute_input":"2025-12-05T12:57:01.130647Z","iopub.status.idle":"2025-12-05T12:57:01.149374Z","shell.execute_reply.started":"2025-12-05T12:57:01.130621Z","shell.execute_reply":"2025-12-05T12:57:01.148641Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"images, labels = next(iter(train_loader))\nprint(\"Batch images shape:\", images.shape)\nprint(\"Batch labels shape:\", labels.shape)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T12:57:01.150036Z","iopub.execute_input":"2025-12-05T12:57:01.150206Z","iopub.status.idle":"2025-12-05T13:00:15.819376Z","shell.execute_reply.started":"2025-12-05T12:57:01.150192Z","shell.execute_reply":"2025-12-05T13:00:15.818657Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.preprocessing import LabelEncoder\n\nlabel_col = \"label\"   # change if your column has a different name\n\nle = LabelEncoder()\ntrain_df[\"label_enc\"] = le.fit_transform(train_df[label_col])\n\nnum_classes = train_df[\"label_enc\"].nunique()\nprint(\"Number of classes:\", num_classes)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T13:00:15.820273Z","iopub.execute_input":"2025-12-05T13:00:15.820511Z","iopub.status.idle":"2025-12-05T13:00:15.827884Z","shell.execute_reply.started":"2025-12-05T13:00:15.820492Z","shell.execute_reply":"2025-12-05T13:00:15.827297Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create the InceptionV3 + Swin hybrid model instead of plain Swin\nmodel = InceptionSwinHybrid(num_classes=num_classes)\nmodel.to(device)\n\n# Quick sanity check\nx = torch.randn(1, 3, 224, 224).to(device)  # IMG_SIZE must be 224\nwith torch.no_grad():\n    out = model(x)\n\nprint(\"Hybrid model output shape:\", out.shape)  # should be [1, num_classes]\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T13:00:15.828655Z","iopub.execute_input":"2025-12-05T13:00:15.828868Z","iopub.status.idle":"2025-12-05T13:00:17.823165Z","shell.execute_reply.started":"2025-12-05T13:00:15.828851Z","shell.execute_reply":"2025-12-05T13:00:17.822345Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\noptimizer = optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4)\n\nEPOCHS = 5  \nscheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T13:00:17.823972Z","iopub.execute_input":"2025-12-05T13:00:17.824227Z","iopub.status.idle":"2025-12-05T13:00:17.83012Z","shell.execute_reply.started":"2025-12-05T13:00:17.824209Z","shell.execute_reply":"2025-12-05T13:00:17.829323Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, criterion, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n\n    for images, labels in loader:\n        images = images.to(device)\n        labels = labels.to(device)\n\n        optimizer.zero_grad()\n\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item() * images.size(0)\n        _, preds = torch.max(outputs, 1)\n        correct += (preds == labels).sum().item()\n        total += labels.size(0)\n\n    return running_loss / total, correct / total\n\n\ndef eval_one_epoch(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n\n    with torch.no_grad():\n        for images, labels in loader:\n            images = images.to(device)\n            labels = labels.to(device)\n\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n\n            running_loss += loss.item() * images.size(0)\n            _, preds = torch.max(outputs, 1)\n            correct += (preds == labels).sum().item()\n            total += labels.size(0)\n\n    return running_loss / total, correct / total\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T13:00:17.830776Z","iopub.execute_input":"2025-12-05T13:00:17.830948Z","iopub.status.idle":"2025-12-05T13:00:17.846754Z","shell.execute_reply.started":"2025-12-05T13:00:17.830934Z","shell.execute_reply":"2025-12-05T13:00:17.846131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_val_acc = 0.0\n\nfor epoch in range(1, EPOCHS + 1):\n    train_loss, train_acc = train_one_epoch(model, train_loader, optimizer, criterion, device)\n    val_loss, val_acc     = eval_one_epoch(model, val_loader, criterion, device)\n\n    scheduler.step()\n\n    print(f\"Epoch {epoch}/{EPOCHS}\")\n    print(f\"  Train loss: {train_loss:.4f}, acc: {train_acc:.4f}\")\n    print(f\"  Val   loss: {val_loss:.4f}, acc: {val_acc:.4f}\")\n\n    if val_acc > best_val_acc:\n        best_val_acc = val_acc\n        torch.save(model.state_dict(), \"best_inception_swin_hybrid.pth\")\n        print(\"Saved new best model (val_acc = {:.4f})\".format(val_acc))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T13:00:17.847586Z","iopub.execute_input":"2025-12-05T13:00:17.847827Z"}},"outputs":[],"execution_count":null}]}