{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":13836,"databundleVersionId":1718836,"sourceType":"competition"}],"dockerImageVersionId":30887,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\n\nimport pandas as pd\nimport torch\nfrom PIL import Image\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision import transforms\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\ntransform = transforms.Compose(\n    [\n        transforms.Resize((448, 448)),  # Resize to input size of MaiaNet\n        transforms.RandomHorizontalFlip(p=0.5),  # Horizontal flipping\n        transforms.RandomVerticalFlip(p=0.5),  # Vertical flipping\n        transforms.ToTensor(),  # Convert to tensor before adding noise\n        transforms.Lambda(lambda x: x + torch.randn_like(x) * 0.05),  # Add Gaussian noise\n        transforms.Lambda(lambda x: transforms.functional.erase(x, i=0, j=0, h=50, w=50, v=0.0)),  # Add cutout\n    ]\n)\n\n# df = pd.read_csv(\"train.csv\")\ndf = pd.read_csv(\"/kaggle/input/cassava-leaf-disease-classification/train.csv\")\n\nprint(df.label.value_counts())\nbalanced_df = pd.DataFrame()\n\nfor label in df[\"label\"].unique():\n    label_df = df[df[\"label\"] == label]\n    if len(label_df) > 1000:\n        _, sampled_df = train_test_split(label_df, test_size=1000, random_state=42, stratify=label_df[\"label\"])\n    balanced_df = pd.concat([balanced_df, sampled_df])\n\n\nclass Dataset(Dataset):\n    def __init__(self, dataframe):\n        self.dataframe = dataframe\n        self.image_dir = \"/kaggle/input/cassava-leaf-disease-classification/train_images\"\n        self.transform = transform\n        self.device = device\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, idx):\n        image_id = self.dataframe.iloc[idx][\"image_id\"]\n        label = self.dataframe.iloc[idx][\"label\"]\n        image_path = os.path.join(self.image_dir, image_id)\n        image = Image.open(image_path).convert(\"RGB\")\n        image = self.transform(image)\n        # Move tensors to GPU if available\n        image = image.to(device)\n        label = torch.tensor(label, device=device)\n        return image, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-22T00:51:50.589740Z","iopub.execute_input":"2025-02-22T00:51:50.590066Z","iopub.status.idle":"2025-02-22T00:51:50.634141Z","shell.execute_reply.started":"2025-02-22T00:51:50.590040Z","shell.execute_reply":"2025-02-22T00:51:50.633157Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df, temp_df = train_test_split(balanced_df, test_size=0.3, random_state=42, stratify=balanced_df[\"label\"])\nval_df, test_df = train_test_split(temp_df, test_size=0.5, random_state=42, stratify=temp_df[\"label\"])\n\ntrain_dataset = Dataset(train_df)\ntest_dataset = Dataset(test_df)\nval_dataset = Dataset(val_df)\n\ntrain_loader = DataLoader(train_dataset, batch_size=24, shuffle=True)\ntest_loader = DataLoader(test_dataset, batch_size=24, shuffle=False)\nval_loader = DataLoader(val_dataset, batch_size=24, shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-22T00:51:53.099428Z","iopub.execute_input":"2025-02-22T00:51:53.099810Z","iopub.status.idle":"2025-02-22T00:51:53.111681Z","shell.execute_reply.started":"2025-02-22T00:51:53.099779Z","shell.execute_reply":"2025-02-22T00:51:53.110857Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.cuda.amp as amp  # For mixed precision training\nimport torch.nn as nn\nimport torch.optim as optim\nfrom sklearn.metrics import accuracy_score, precision_recall_fscore_support\nfrom torch.optim.lr_scheduler import ExponentialLR\nfrom tqdm import tqdm\nfrom torch.amp import GradScaler, autocast\n\n\nclass Trainer:\n    def __init__(self, model, train_loader, val_loader, test_loader, lr=0.2, num_epochs=80):\n        self.device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n        self.model = model.to(self.device)\n\n        # Enable cudnn benchmarking for better performance\n        if torch.cuda.is_available():\n            torch.backends.cudnn.benchmark = True\n\n        self.train_loader = train_loader\n        self.val_loader = val_loader\n        self.test_loader = test_loader\n        self.num_epochs = num_epochs\n        self.lr = lr\n\n        self.optimizer = optim.SGD(self.model.parameters(), lr=self.lr, momentum=0.9, weight_decay=1e-5)\n        self.scheduler = ExponentialLR(self.optimizer, gamma=0.96)\n        self.criterion = nn.CrossEntropyLoss().to(self.device)  # Move loss function to GPU\n\n        # Initialize mixed precision training\n        self.scaler = torch.amp.GradScaler('cuda')\n\n        self.best_val_loss = float(\"inf\")\n        self.best_model_state = None\n\n    def train_epoch(self, epoch):\n        self.model.train()\n        total_loss = 0\n\n        pbar = tqdm(self.train_loader, desc=f\"Epoch {epoch + 1}/{self.num_epochs}\")\n\n        for images, labels in pbar:\n            # Clear GPU cache if needed\n            if torch.cuda.is_available():\n                torch.cuda.empty_cache()\n\n            images = images.to(self.device, non_blocking=True)\n            labels = labels.to(self.device, non_blocking=True)\n\n            self.optimizer.zero_grad(set_to_none=True)  # More efficient than zero_grad()\n\n            # Use mixed precision training\n            with amp.autocast('cuda'):\n                outputs = self.model(images)\n                loss = self.criterion(outputs, labels)\n\n            # Scale the loss and perform backprop\n            self.scaler.scale(loss).backward()\n            self.scaler.step(self.optimizer)\n            self.scaler.update()\n\n            total_loss += loss.item()\n            pbar.set_postfix(loss=f\"{loss.item():.4f}\")\n\n            # Delete unnecessary tensors\n            del outputs, loss\n\n        self.scheduler.step()\n\n        return total_loss / len(self.train_loader)\n\n    @torch.no_grad()  # More efficient than with torch.no_grad()\n    def validate(self):\n        self.model.eval()\n        total_loss = 0\n        all_preds, all_labels = [], []\n\n        for images, labels in self.val_loader:\n            images = images.to(self.device, non_blocking=True)\n            labels = labels.to(self.device, non_blocking=True)\n\n            with amp.autocast('cuda'):\n                outputs = self.model(images)\n                loss = self.criterion(outputs, labels)\n\n            total_loss += loss.item() * labels.size(0)\n\n            preds = torch.argmax(outputs, dim=1)\n            all_preds.extend(preds.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n\n            # Clean up GPU memory\n            del outputs, loss\n\n        avg_loss = total_loss / len(self.val_loader.dataset)\n        metrics = self.calculate_metrics(all_preds, all_labels)\n\n        return avg_loss, metrics\n\n    @staticmethod\n    def calculate_metrics(predictions, labels):\n        accuracy = accuracy_score(labels, predictions)\n        precision, recall, f1, _ = precision_recall_fscore_support(labels, predictions, average=\"weighted\", zero_division=0)\n        return {\"accuracy\": accuracy, \"precision\": precision, \"recall\": recall, \"f1\": f1}\n\n    @staticmethod\n    def print_metrics(metrics, phase):\n        print(f\"\\n{phase} Metrics:\")\n        print(\"-\" * 50)\n        for metric, value in metrics.items():\n            print(f\"{metric.capitalize()}: {value:.4f}\")\n        print(\"-\" * 50)\n\n    def train(self):\n        try:\n            for epoch in range(self.num_epochs):\n                train_loss = self.train_epoch(epoch)\n                val_loss, val_metrics = self.validate()\n\n                print(f\"\\nEpoch {epoch + 1}: Train Loss = {train_loss:.4f} | Val Loss = {val_loss:.4f}\")\n                self.print_metrics(val_metrics, \"Validation\")\n\n                if val_loss < self.best_val_loss:\n                    self.best_val_loss = val_loss\n                    # Save model state to CPU to avoid GPU memory issues\n                    self.best_model_state = {k: v.cpu() for k, v in self.model.state_dict().items()}\n\n                # Print GPU memory usage if available\n                if torch.cuda.is_available():\n                    print(f\"GPU Memory allocated: {torch.cuda.memory_allocated() / 1e9:.2f} GB\")\n\n        except Exception as e:\n            print(f\"Training interrupted: {str(e)}\")\n            # Save the current best model if training is interrupted\n            if self.best_model_state is not None:\n                torch.save(self.best_model_state, \"interrupted_model.pt\")\n\n    def test(self):\n        # Load best model state back to GPU\n        if self.best_model_state is not None:\n            self.model.load_state_dict({k: v.to(self.device) for k, v in self.best_model_state.items()})\n        test_loss, test_metrics = self.validate()\n        print(\"\\nBest Model Performance on Test Set:\")\n        self.print_metrics(test_metrics, \"Test\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-22T00:57:22.253530Z","iopub.execute_input":"2025-02-22T00:57:22.253879Z","iopub.status.idle":"2025-02-22T00:57:22.269230Z","shell.execute_reply.started":"2025-02-22T00:57:22.253844Z","shell.execute_reply":"2025-02-22T00:57:22.268265Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model code","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nimport math\n\nimport torch\nimport torch.nn as nn\nfrom torchvision.models import ResNet\n\n\ndef get_freq_indices(method):\n    assert method in [\"top1\", \"top2\", \"top4\", \"top8\", \"top16\", \"top32\", \"bot1\", \"bot2\", \"bot4\", \"bot8\", \"bot16\", \"bot32\", \"low1\", \"low2\", \"low4\", \"low8\", \"low16\", \"low32\"]\n    num_freq = int(method[3:])\n    if \"top\" in method:\n        all_top_indices_x = [0, 0, 6, 0, 0, 1, 1, 4, 5, 1, 3, 0, 0, 0, 3, 2, 4, 6, 3, 5, 5, 2, 6, 5, 5, 3, 3, 4, 2, 2, 6, 1]\n        all_top_indices_y = [0, 1, 0, 5, 2, 0, 2, 0, 0, 6, 0, 4, 6, 3, 5, 2, 6, 3, 3, 3, 5, 1, 1, 2, 4, 2, 1, 1, 3, 0, 5, 3]\n        mapper_x = all_top_indices_x[:num_freq]\n        mapper_y = all_top_indices_y[:num_freq]\n    elif \"low\" in method:\n        all_low_indices_x = [0, 0, 1, 1, 0, 2, 2, 1, 2, 0, 3, 4, 0, 1, 3, 0, 1, 2, 3, 4, 5, 0, 1, 2, 3, 4, 5, 6, 1, 2, 3, 4]\n        all_low_indices_y = [0, 1, 0, 1, 2, 0, 1, 2, 2, 3, 0, 0, 4, 3, 1, 5, 4, 3, 2, 1, 0, 6, 5, 4, 3, 2, 1, 0, 6, 5, 4, 3]\n        mapper_x = all_low_indices_x[:num_freq]\n        mapper_y = all_low_indices_y[:num_freq]\n    elif \"bot\" in method:\n        all_bot_indices_x = [6, 1, 3, 3, 2, 4, 1, 2, 4, 4, 5, 1, 4, 6, 2, 5, 6, 1, 6, 2, 2, 4, 3, 3, 5, 5, 6, 2, 5, 5, 3, 6]\n        all_bot_indices_y = [6, 4, 4, 6, 6, 3, 1, 4, 4, 5, 6, 5, 2, 2, 5, 1, 4, 3, 5, 0, 3, 1, 1, 2, 4, 2, 1, 1, 5, 3, 3, 3]\n        mapper_x = all_bot_indices_x[:num_freq]\n        mapper_y = all_bot_indices_y[:num_freq]\n    else:\n        raise NotImplementedError\n    return mapper_x, mapper_y\n\n\nclass MultiSpectralAttentionLayer(torch.nn.Module):\n    def __init__(self, channel, dct_h, dct_w, reduction=16, freq_sel_method=\"top16\"):\n        super(MultiSpectralAttentionLayer, self).__init__()\n        self.reduction = reduction\n        self.dct_h = dct_h\n        self.dct_w = dct_w\n\n        mapper_x, mapper_y = get_freq_indices(freq_sel_method)\n        self.num_split = len(mapper_x)\n        mapper_x = [temp_x * (dct_h // 7) for temp_x in mapper_x]\n        mapper_y = [temp_y * (dct_w // 7) for temp_y in mapper_y]\n        # make the frequencies in different sizes are identical to a 7x7 frequency space\n        # eg, (2,2) in 14x14 is identical to (1,1) in 7x7\n\n        self.dct_layer = MultiSpectralDCTLayer(dct_h, dct_w, mapper_x, mapper_y, channel)\n        self.fc = nn.Sequential(nn.Linear(channel, channel // reduction, bias=False), nn.ReLU(), nn.Linear(channel // reduction, channel, bias=False), nn.Sigmoid())\n\n    def forward(self, x):\n        n, c, h, w = x.shape\n        x_pooled = x\n        if h != self.dct_h or w != self.dct_w:\n            x_pooled = torch.nn.functional.adaptive_avg_pool2d(x, (self.dct_h, self.dct_w))\n            # If you have concerns about one-line-change, don't worry.   :)\n            # In the ImageNet models, this line will never be triggered.\n            # This is for compatibility in instance segmentation and object detection.\n        y = self.dct_layer(x_pooled)\n\n        y = self.fc(y).view(n, c, 1, 1)\n        return x * y.expand_as(x)\n\n\nclass MultiSpectralDCTLayer(nn.Module):\n    \"\"\"\n    Generate dct filters\n    \"\"\"\n\n    def __init__(self, height, width, mapper_x, mapper_y, channel):\n        super(MultiSpectralDCTLayer, self).__init__()\n\n        assert len(mapper_x) == len(mapper_y)\n        assert channel % len(mapper_x) == 0\n\n        self.num_freq = len(mapper_x)\n\n        # fixed DCT init\n        self.register_buffer(\"weight\", self.get_dct_filter(height, width, mapper_x, mapper_y, channel))\n\n        # fixed random init\n        # self.register_buffer('weight', torch.rand(channel, height, width))\n\n        # learnable DCT init\n        # self.register_parameter('weight', self.get_dct_filter(height, width, mapper_x, mapper_y, channel))\n\n        # learnable random init\n        # self.register_parameter('weight', torch.rand(channel, height, width))\n\n        # num_freq, h, w\n\n    def forward(self, x):\n        assert len(x.shape) == 4, \"x must been 4 dimensions, but got \" + str(len(x.shape))\n        # n, c, h, w = x.shape\n\n        x = x * self.weight\n\n        result = torch.sum(x, dim=[2, 3])\n        return result\n\n    def build_filter(self, pos, freq, POS):\n        result = math.cos(math.pi * freq * (pos + 0.5) / POS) / math.sqrt(POS)\n        if freq == 0:\n            return result\n        else:\n            return result * math.sqrt(2)\n\n    def get_dct_filter(self, tile_size_x, tile_size_y, mapper_x, mapper_y, channel):\n        dct_filter = torch.zeros(channel, tile_size_x, tile_size_y)\n\n        c_part = channel // len(mapper_x)\n\n        for i, (u_x, v_y) in enumerate(zip(mapper_x, mapper_y)):\n            for t_x in range(tile_size_x):\n                for t_y in range(tile_size_y):\n                    dct_filter[i * c_part : (i + 1) * c_part, t_x, t_y] = self.build_filter(t_x, u_x, tile_size_x) * self.build_filter(t_y, v_y, tile_size_y)\n\n        return dct_filter\n\n\ndef conv3x3(in_planes, out_planes, stride=1):\n    return nn.Conv2d(in_planes, out_planes, kernel_size=3, stride=stride, padding=1, bias=False)\n    \nimport torch.nn as nn\nfrom torch.hub import load_state_dict_from_url\nfrom torchvision.models import ResNet\n\n\nclass SELayer(nn.Module):\n    def __init__(self, channel, reduction=16):\n        super(SELayer, self).__init__()\n        self.avg_pool = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Sequential(nn.Linear(channel, channel // reduction, bias=False), nn.ReLU(), nn.Linear(channel // reduction, channel, bias=False), nn.Sigmoid())\n\n    def forward(self, x):\n        b, c, _, _ = x.size()\n        y = self.avg_pool(x).view(b, c)\n        y = self.fc(y).view(b, c, 1, 1)\n        return x * y.expand_as(x)\n\n\ndef conv3x3(in_planes, out_planes, stride=1):\n    return nn.Conv2d(in_planes, out_planes, kernel_size=3, stride=stride, padding=1, bias=False)\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n\nclass MaiaNet(nn.Module):\n    def __init__(self, num_classes):\n        super(MaiaNet, self).__init__()\n        self.head = HeadBlock(3, 64)  # Input: 448×448×3 -> 112×112×64\n        self.anti_aliasing_1 = AntiAliasingBlock(64, 64, downsample=False)  # 112×112×64 -> 112×112×64\n        self.maia_1 = MaiaBlock(64, 256)  # 112×112×64 -> 112×112×256\n        self.anti_aliasing_2 = AntiAliasingBlock(256, 512, downsample=True)  # 112×112×256 -> 56×56×512\n        self.maia_2 = MaiaBlock(512, 512)  # 56×56×512 -> 56×56×512\n        self.anti_aliasing_3 = AntiAliasingBlock(512, 1024, downsample=True)  # 56×56×512 -> 28×28×1024\n        self.maia_3 = MaiaBlock(1024, 1024)  # 28×28×1024 -> 28×28×1024\n        self.maia_4 = MaiaBlock(1024, 2048, downsample=True)  # 14×14×2048 -> 14×14×2048\n        self.global_pool = nn.AdaptiveAvgPool2d((1, 1))  # Converts 14x14x2048 to 1x1x2048\n        self.fc = nn.Linear(2048, num_classes)  # Fully connected layer (2048 -> num_classes)\n        self.softmax = nn.Softmax(dim=1)\n\n    def forward(self, x, verbose=False):\n        if verbose:\n            print(\"Input:\", x.shape)\n        x = self.head(x)\n        if verbose:\n            print(\"Head:\", x.shape)\n        x = self.anti_aliasing_1(x)\n        if verbose:\n            print(\"Anti-aliasing 1:\", x.shape)\n        x = self.maia_1(x)\n        if verbose:\n            print(\"MAIA 1:\", x.shape)\n        x = self.anti_aliasing_2(x)\n        if verbose:\n            print(\"Anti-aliasing 2:\", x.shape)\n        x = self.maia_2(x)\n        if verbose:\n            print(\"MAIA 2:\", x.shape)\n        x = self.anti_aliasing_3(x)\n        if verbose:\n            print(\"Anti-aliasing 3:\", x.shape)\n        x = self.maia_3(x)\n        if verbose:\n            print(\"MAIA 3:\", x.shape)\n        x = self.maia_4(x)\n        if verbose:\n            print(\"MAIA 4:\", x.shape)\n\n        x = self.global_pool(x)  # Shape: (batch_size, 2048, 1, 1)\n        x = torch.flatten(x, 1)  # Shape: (batch_size, 2048)\n        x = self.fc(x)  # Shape: (batch_size, num_classes)\n\n        return x\n\n\nclass HeadBlock(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(HeadBlock, self).__init__()\n        self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=7, padding=3, stride=2)\n        self.bn = nn.BatchNorm2d(out_channels)\n        self.pool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)\n\n    def forward(self, x):\n        x = self.conv(x)\n        x = self.bn(x)\n        x = F.relu(x)\n        x = self.pool(x)\n        return x\n\n\nclass MultiAttention(nn.Module):\n    def __init__(self, in_channels):\n        super(MultiAttention, self).__init__()\n\n        # https://github.com/hujie-frank/SENet/blob/master/README.md\n        self.se = SELayer(in_channels, reduction=16)\n\n        # https://github.com/cfzd/FcaNet/blob/master/model/fcanet.py\n        self.fca = MultiSpectralAttentionLayer(in_channels, 7, 7, reduction=16, freq_sel_method=\"top16\")\n\n    def forward(self, x):\n        x = self.se(x)\n        x = self.fca(x)\n        return x\n\n\nclass AntiAliasingBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, downsample=True):\n        super(AntiAliasingBlock, self).__init__()\n\n        self.downsample = downsample\n\n        self.block1 = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(),\n        )\n\n        self.down_conversion = nn.Sequential(\n            nn.SiLU(),\n            nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1, groups=out_channels),\n            nn.BatchNorm2d(out_channels),\n            nn.SiLU(),\n        )\n\n        stride = 2 if self.downsample else 1\n        self.block2 = nn.Sequential(\n            nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=stride, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(),\n            nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),\n        )\n\n        self.ma = MultiAttention(out_channels)\n        self.ibn = nn.InstanceNorm2d(out_channels)\n\n        self.residual_conv = None\n        if in_channels != out_channels or downsample:\n            self.residual_conv = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride)\n        else:\n            self.residual_conv = None\n\n    def forward(self, x):\n        out = self.block1(x)\n        out = self.down_conversion(out)\n        out = self.block2(out)\n        out = self.ma(out)\n        if self.residual_conv:\n            x = self.residual_conv(x)\n        out = out + x\n        out = self.ibn(out)\n        out = F.relu(out)\n        return out\n\n\nclass MaiaBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, downsample=False):\n        super(MaiaBlock, self).__init__()\n\n        stride = 2 if downsample else 1\n\n        self.conv1 = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1, stride=stride),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(),\n        )\n\n        self.conv2 = nn.Sequential(\n            nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1, stride=1),\n            nn.BatchNorm2d(out_channels),\n            nn.SiLU(),\n        )\n\n        self.conv3 = nn.Sequential(\n            nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1, stride=1),\n            nn.BatchNorm2d(out_channels),\n        )\n\n        self.ma = MultiAttention(out_channels)\n        self.ibn = nn.InstanceNorm2d(out_channels)\n\n        if in_channels != out_channels or downsample:\n            self.residual_conv = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride)\n        else:\n            self.residual_conv = None\n\n    def forward(self, x):\n        out = self.conv1(x)\n        out = self.conv2(out)\n        out = self.conv3(out)\n        out = self.ma(out)\n\n        if self.residual_conv:\n            x = self.residual_conv(x)\n        out = out + x\n        out = self.ibn(out)\n        out = F.relu(out)\n        return out\n\n\nif __name__ == \"__main__\":\n    model = MaiaNet(num_classes=5).to(device)\n    x = torch.randn(1, 3, 448, 448).to(device)\n    output = model(x)\n    print(output.shape)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-22T00:52:06.774462Z","iopub.execute_input":"2025-02-22T00:52:06.774808Z","iopub.status.idle":"2025-02-22T00:52:08.552563Z","shell.execute_reply.started":"2025-02-22T00:52:06.774780Z","shell.execute_reply":"2025-02-22T00:52:08.551794Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.cuda.amp import GradScaler, autocast\nfrom sklearn.metrics import accuracy_score, precision_recall_fscore_support\nfrom torch.optim.lr_scheduler import OneCycleLR\nfrom tqdm import tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-22T00:56:24.901410Z","iopub.execute_input":"2025-02-22T00:56:24.901820Z","iopub.status.idle":"2025-02-22T00:56:24.906650Z","shell.execute_reply.started":"2025-02-22T00:56:24.901785Z","shell.execute_reply":"2025-02-22T00:56:24.905625Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer = Trainer(model, train_loader, test_loader, val_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-22T00:57:25.461229Z","iopub.execute_input":"2025-02-22T00:57:25.461534Z","iopub.status.idle":"2025-02-22T00:57:25.469022Z","shell.execute_reply.started":"2025-02-22T00:57:25.461510Z","shell.execute_reply":"2025-02-22T00:57:25.468166Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer.train()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-22T00:57:26.901514Z","iopub.execute_input":"2025-02-22T00:57:26.901832Z","execution_failed":"2025-02-22T02:23:37.754Z"}},"outputs":[],"execution_count":null}]}