{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":[{"sourceId":21154,"databundleVersionId":1243559,"sourceType":"competition"}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"pip install tfrecord","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-29T20:03:26.556991Z","iopub.execute_input":"2025-12-29T20:03:26.557538Z","iopub.status.idle":"2025-12-29T20:03:35.256303Z","shell.execute_reply.started":"2025-12-29T20:03:26.557507Z","shell.execute_reply":"2025-12-29T20:03:35.255440Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import io\nimport math\nimport os\nimport random\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom PIL import Image\n\nimport torch\nfrom torch import nn\nfrom torch.utils.data import IterableDataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nfrom torchvision import transforms\n\nfrom sklearn.metrics import f1_score, precision_score, recall_score, confusion_matrix\n\nimport timm\n\ntry:\n    import tfrecord\nexcept ModuleNotFoundError:\n    tfrecord = None\n    print(\"Warning: tfrecord not installed.\")\n\nfrom tqdm.auto import tqdm\n\nplt.style.use('seaborn-v0_8')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-29T20:03:35.258095Z","iopub.execute_input":"2025-12-29T20:03:35.258348Z","iopub.status.idle":"2025-12-29T20:03:53.123120Z","shell.execute_reply.started":"2025-12-29T20:03:35.258315Z","shell.execute_reply":"2025-12-29T20:03:53.122438Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ModelConfig:\n    base_path = Path('/kaggle/input/tpu-getting-started')\n    tfrecord_dir = base_path / 'tfrecords-jpeg-512x512'\n    train_dir = tfrecord_dir / 'train'\n    val_dir = tfrecord_dir / 'val'\n    test_dir = tfrecord_dir / 'test'\n\n    image_size = (384, 384)\n    num_classes = 104\n\n    batch_size = 16\n    val_batch_size = 64\n    test_batch_size = 64\n    epochs = 10\n\n    learning_rate = 3e-4\n    min_lr = 1e-6\n    weight_decay = 1e-5\n    grad_clip = 1.0\n\n    num_workers = 1\n    seed = 2235\n\n    model_name = 'swin_large_patch4_window12_384'\n    drop_rate = 0.214\n    drop_path_rate = 0.118\n    use_amp = True\n\n    output_dir = Path('/kaggle/working')\n    model_path = output_dir / 'best_model.pt'\n    submission_path = output_dir / 'submission.csv'\n\n    train_limit = None\n    val_limit = None\n    test_limit = None\n\n    classes = [\n        'pink primrose', 'hard-leaved pocket orchid', 'canterbury bells', 'sweet pea', 'wild geranium',\n        'tiger lily', 'moon orchid', 'bird of paradise', 'monkshood', 'globe thistle',\n        'snapdragon', \"colt's foot\", 'king protea', 'spear thistle', 'yellow iris',\n        'globe-flower', 'purple coneflower', 'peruvian lily', 'balloon flower', 'giant white arum lily',\n        'fire lily', 'pincushion flower', 'fritillary', 'red ginger', 'grape hyacinth',\n        'corn poppy', 'prince of wales feathers', 'stemless gentian', 'artichoke', 'sweet william',\n        'carnation', 'garden phlox', 'love in the mist', 'cosmos', 'alpine sea holly',\n        'ruby-lipped cattleya', 'cape flower', 'great masterwort', 'siam tulip', 'lenten rose',\n        'barberton daisy', 'daffodil', 'sword lily', 'poinsettia', 'bolero deep blue',\n        'wallflower', 'marigold', 'buttercup', 'daisy', 'common dandelion',\n        'petunia', 'wild pansy', 'primula', 'sunflower', 'lilac hibiscus',\n        'bishop of llandaff', 'gaura', 'geranium', 'orange dahlia', 'pink-yellow dahlia',\n        'cautleya spicata', 'japanese anemone', 'black-eyed susan', 'silverbush', 'californian poppy',\n        'osteospermum', 'spring crocus', 'iris', 'windflower', 'tree poppy',\n        'gazania', 'azalea', 'water lily', 'rose', 'thorn apple',\n        'morning glory', 'passion flower', 'lotus', 'toad lily', 'anthurium',\n        'frangipani', 'clematis', 'hibiscus', 'columbine', 'desert-rose',\n        'tree mallow', 'magnolia', 'cyclamen ', 'watercress', 'canna lily',\n        'hippeastrum ', 'bee balm', 'pink quill', 'foxglove', 'bougainvillea',\n        'camellia', 'mallow', 'mexican petunia', 'bromelia', 'blanket flower',\n        'trumpet creeper', 'blackberry lily', 'common tulip', 'wild rose'\n    ]\n\n\nModelConfig.output_dir.mkdir(parents=True, exist_ok=True)\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n\ndef set_seed(seed: int = ModelConfig.seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\n\nset_seed()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-29T20:03:53.124039Z","iopub.execute_input":"2025-12-29T20:03:53.124262Z","iopub.status.idle":"2025-12-29T20:03:53.234562Z","shell.execute_reply.started":"2025-12-29T20:03:53.124239Z","shell.execute_reply":"2025-12-29T20:03:53.233792Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"use_cuda = torch.cuda.is_available()\nnum_gpus = torch.cuda.device_count() if use_cuda else 0\ndevice = torch.device('cuda' if use_cuda else 'cpu')\n\n\ndef set_seed(seed: int = ModelConfig.seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if use_cuda:\n        torch.cuda.manual_seed_all(seed)\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False\n\n\nset_seed()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-29T20:03:53.236136Z","iopub.execute_input":"2025-12-29T20:03:53.236362Z","iopub.status.idle":"2025-12-29T20:03:53.280710Z","shell.execute_reply.started":"2025-12-29T20:03:53.236340Z","shell.execute_reply":"2025-12-29T20:03:53.280152Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AverageMeter:\n\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0.0\n        self.avg = 0.0\n        self.sum = 0.0\n        self.count = 0\n\n    def update(self, val: float, n: int = 1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count if self.count else 0.0\n\n\ndef accuracy_from_logits(logits: torch.Tensor, targets: torch.Tensor) -> float:\n    preds = logits.argmax(dim=1)\n    return (preds == targets).float().mean().item()\n\n\ndef compute_classification_metrics(preds: np.ndarray, labels: np.ndarray):\n    accuracy = (preds == labels).mean()\n    f1 = f1_score(labels, preds, labels=np.arange(ModelConfig.num_classes), average='macro')\n    precision = precision_score(labels, preds, labels=np.arange(ModelConfig.num_classes), average='macro')\n    recall = recall_score(labels, preds, labels=np.arange(ModelConfig.num_classes), average='macro')\n    cmat = confusion_matrix(labels, preds, labels=np.arange(ModelConfig.num_classes))\n    return accuracy, f1, precision, recall, cmat\n\n\ndef plot_training_curves(history):\n    fig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\n    axes[0].plot(history['train_loss'], label='train')\n    axes[0].plot(history['val_loss'], label='val')\n    axes[0].set_title('Loss')\n    axes[0].set_xlabel('epoch')\n    axes[0].legend()\n\n    axes[1].plot(history['train_acc'], label='train')\n    axes[1].plot(history['val_acc'], label='val')\n    axes[1].set_title('Accuracy')\n    axes[1].set_xlabel('epoch')\n    axes[1].legend()\n\n    axes[2].plot(history['lr'])\n    axes[2].set_title('Learning Rate')\n    axes[2].set_xlabel('epoch')\n\n    plt.show()\n\n\ndef show_confusion_matrix(cmat: np.ndarray, title_suffix: str = ''):\n    fig, ax = plt.subplots(figsize=(14, 12))\n    cax = ax.imshow(cmat, interpolation='nearest', cmap='Reds')\n    fig.colorbar(cax)\n    ax.set_title(f'Confusion Matrix {title_suffix}')\n    ax.set_xlabel('Predicted')\n    ax.set_ylabel('True')\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-29T20:03:53.281432Z","iopub.execute_input":"2025-12-29T20:03:53.281730Z","iopub.status.idle":"2025-12-29T20:03:53.291348Z","shell.execute_reply.started":"2025-12-29T20:03:53.281706Z","shell.execute_reply":"2025-12-29T20:03:53.290797Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import re\n\n\ndef get_tfrecord_files(directory: Path):\n    files = sorted(str(p) for p in directory.glob('*.tfrec'))\n    if not files:\n        raise FileNotFoundError(f'No TFRecord files found in {directory}')\n    return files\n\n\ndef count_data_items(filenames):\n    pattern = re.compile(r'-([0-9]*)\\.tfrec$')\n    total = 0\n    for fname in filenames:\n        match = pattern.search(fname)\n        if match:\n            total += int(match.group(1))\n    return total\n\n\ndef get_train_transforms():\n    return transforms.Compose([\n        transforms.RandomResizedCrop(ModelConfig.image_size, scale=(0.7, 1.0), ratio=(0.8, 1.2)),\n        transforms.RandomHorizontalFlip(p=0.618),\n        transforms.RandomVerticalFlip(p=0.314),\n        transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.05),\n        transforms.RandomRotation(degrees=20),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n        transforms.RandomErasing(p=0.2, scale=(0.02, 0.1), ratio=(0.3, 3.3), value='random'),\n    ])\n\n\ndef get_eval_transforms():\n    return transforms.Compose([\n        transforms.Resize(ModelConfig.image_size),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n    ])\n\n\nclass FlowerTFRecordDataset(IterableDataset):\n    def __init__(self, filenames, labeled=True, transform=None, seed=ModelConfig.seed, limit=None):\n        self.filenames = list(filenames)\n        self.labeled = labeled\n        self.transform = transform\n        self.seed = seed\n        self.limit = limit\n        self.epoch = 0\n        total_items = count_data_items(self.filenames)\n        self.num_samples = min(total_items, limit) if limit is not None else total_items\n\n    def set_epoch(self, epoch: int):\n        self.epoch = epoch\n\n    def __len__(self):\n        return self.num_samples\n\n    def _shuffle_files(self):\n        rng = random.Random(self.seed + self.epoch)\n        files = self.filenames.copy()\n        rng.shuffle(files)\n        return files\n\n    def _shuffle_records(self, records, *, fname):\n        rng = random.Random(self.seed + self.epoch + hash(fname) % 10_000)\n        records = list(records)\n        rng.shuffle(records)\n        return records\n\n    def __iter__(self):\n        files = self._shuffle_files()\n        worker_info = torch.utils.data.get_worker_info()\n        if worker_info is not None:\n            files = files[worker_info.id :: worker_info.num_workers]\n\n        produced = 0\n        for fname in files:\n            if self.limit is not None and produced >= self.limit:\n                break\n\n            try:\n                records = tfrecord.tfrecord_loader(fname, None, None)\n            except Exception as exc:\n                print(f'Warning: failed to read {fname}: {exc}')\n                continue\n\n            for record in self._shuffle_records(records, fname=fname):\n                if self.limit is not None and produced >= self.limit:\n                    break\n\n                image_bytes = record['image']\n                image = Image.open(io.BytesIO(image_bytes)).convert('RGB')\n                if self.transform is not None:\n                    image = self.transform(image)\n\n                if self.labeled:\n                    label_raw = record['class']\n                    if isinstance(label_raw, np.ndarray):\n                        label = int(label_raw.reshape(-1)[0])\n                    elif isinstance(label_raw, (list, tuple)):\n                        label = int(label_raw[0])\n                    else:\n                        label = int(label_raw)\n                    yield image, label\n                else:\n                    image_id = record.get('id', '')\n                    if isinstance(image_id, bytes):\n                        image_id = image_id.decode()\n                    yield image, image_id\n\n                produced += 1\n\n\ndef build_loader(directory: Path, *, labeled: bool, transform, batch_size: int, limit=None):\n    filenames = get_tfrecord_files(directory)\n    dataset = FlowerTFRecordDataset(filenames, labeled=labeled, transform=transform, limit=limit)\n\n    loader = DataLoader(\n        dataset,\n        batch_size=batch_size,\n        shuffle=False,\n        num_workers=ModelConfig.num_workers,\n        pin_memory=True,\n        drop_last=False,\n    )\n\n    return loader, dataset\n\n\ndef get_loaders():\n    train_loader, _ = build_loader(\n        ModelConfig.train_dir,\n        labeled=True,\n        transform=get_train_transforms(),\n        batch_size=ModelConfig.batch_size,\n        limit=ModelConfig.train_limit,\n    )\n\n    val_loader, _ = build_loader(\n        ModelConfig.val_dir,\n        labeled=True,\n        transform=get_eval_transforms(),\n        batch_size=ModelConfig.val_batch_size,\n        limit=ModelConfig.val_limit,\n    )\n\n    return train_loader, val_loader\n\n\ndef get_test_loader():\n    test_loader, _ = build_loader(\n        ModelConfig.test_dir,\n        labeled=False,\n        transform=get_eval_transforms(),\n        batch_size=ModelConfig.test_batch_size,\n        limit=ModelConfig.test_limit,\n    )\n    return test_loader\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-29T20:03:53.292205Z","iopub.execute_input":"2025-12-29T20:03:53.293002Z","iopub.status.idle":"2025-12-29T20:03:53.319729Z","shell.execute_reply.started":"2025-12-29T20:03:53.292962Z","shell.execute_reply":"2025-12-29T20:03:53.319005Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_model():\n    model = timm.create_model(\n        ModelConfig.model_name,\n        pretrained=True,\n        num_classes=ModelConfig.num_classes,\n        drop_rate=ModelConfig.drop_rate,\n        drop_path_rate=ModelConfig.drop_path_rate,\n    )\n    return model\n\n\ndef create_optimizer(model):\n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=ModelConfig.learning_rate,\n        weight_decay=ModelConfig.weight_decay,\n    )\n    return optimizer\n\n\ndef create_scheduler(optimizer):\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer,\n        T_max=ModelConfig.epochs,\n        eta_min=ModelConfig.min_lr,\n    )\n    return scheduler\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-29T20:03:53.320640Z","iopub.execute_input":"2025-12-29T20:03:53.320984Z","iopub.status.idle":"2025-12-29T20:03:53.339779Z","shell.execute_reply.started":"2025-12-29T20:03:53.320950Z","shell.execute_reply":"2025-12-29T20:03:53.339200Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_epoch(loader, model, criterion, optimizer=None, scaler=None, epoch=0, training=True):\n    if training and hasattr(loader.dataset, 'set_epoch'):\n        loader.dataset.set_epoch(epoch)\n\n    if training:\n        model.train()\n    else:\n        model.eval()\n\n    losses = AverageMeter()\n    accuracies = AverageMeter()\n\n    total_steps = math.ceil(len(loader.dataset) / loader.batch_size)\n\n    pbar = tqdm(loader, total=total_steps, leave=False)\n    pbar.set_description(f\"Epoch {epoch + 1} [{'train' if training else 'val'}]\")\n\n    for images, targets in pbar:\n        images = images.to(device, non_blocking=True)\n        targets = targets.to(device, non_blocking=True)\n\n        with torch.set_grad_enabled(training):\n            with autocast(enabled=scaler is not None and scaler.is_enabled()):\n                outputs = model(images)\n                loss = criterion(outputs, targets)\n\n        batch_size = targets.size(0)\n        if training and optimizer is not None:\n            optimizer.zero_grad(set_to_none=True)\n            if scaler is not None and scaler.is_enabled():\n                scaler.scale(loss).backward()\n                torch.nn.utils.clip_grad_norm_(model.parameters(), ModelConfig.grad_clip)\n                scaler.step(optimizer)\n                scaler.update()\n            else:\n                loss.backward()\n                torch.nn.utils.clip_grad_norm_(model.parameters(), ModelConfig.grad_clip)\n                optimizer.step()\n\n        acc = accuracy_from_logits(outputs.detach(), targets)\n        losses.update(loss.item(), batch_size)\n        accuracies.update(acc, batch_size)\n\n        pbar.set_postfix({\n            'loss': f'{losses.avg:.4f}',\n            'acc': f'{accuracies.avg:.4f}',\n            'lr': optimizer.param_groups[0]['lr'] if optimizer is not None else 0.0,\n        })\n\n    return losses.avg, accuracies.avg\n\n\ndef train_model():\n    train_loader, val_loader = get_loaders()\n\n    model = build_model().to(device)\n    if use_cuda and num_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f'Using DataParallel on {num_gpus} GPUs')\n    criterion = nn.CrossEntropyLoss()\n    optimizer = create_optimizer(model)\n    scheduler = create_scheduler(optimizer)\n    scaler = GradScaler(enabled=ModelConfig.use_amp and device.type == 'cuda')\n\n    history = {\n        'train_loss': [],\n        'train_acc': [],\n        'val_loss': [],\n        'val_acc': [],\n        'lr': [],\n    }\n\n    best_val_acc = 0.0\n\n    for epoch in range(ModelConfig.epochs):\n        train_loss, train_acc = run_epoch(\n            train_loader,\n            model,\n            criterion,\n            optimizer=optimizer,\n            scaler=scaler,\n            epoch=epoch,\n            training=True,\n        )\n\n        val_loss, val_acc = run_epoch(\n            val_loader,\n            model,\n            criterion,\n            optimizer=None,\n            scaler=None,\n            epoch=epoch,\n            training=False,\n        )\n\n        scheduler.step()\n        current_lr = optimizer.param_groups[0]['lr']\n\n        history['train_loss'].append(train_loss)\n        history['train_acc'].append(train_acc)\n        history['val_loss'].append(val_loss)\n        history['val_acc'].append(val_acc)\n        history['lr'].append(current_lr)\n\n        print(f\"Epoch {epoch + 1}/{ModelConfig.epochs}: train_loss={train_loss:.4f}, train_acc={train_acc:.4f}, val_loss={val_loss:.4f}, val_acc={val_acc:.4f}, lr={current_lr:.2e}\")\n\n        if val_acc > best_val_acc:\n            best_val_acc = val_acc\n            model_to_save = model.module if isinstance(model, nn.DataParallel) else model\n            torch.save({'model_state_dict': model_to_save.state_dict(), 'val_acc': best_val_acc}, ModelConfig.model_path)\n            print(f\"  ✓ New best model saved (val_acc={best_val_acc:.4f})\")\n\n    return model, history, best_val_acc\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-29T20:03:53.340651Z","iopub.execute_input":"2025-12-29T20:03:53.341220Z","iopub.status.idle":"2025-12-29T20:03:53.359109Z","shell.execute_reply.started":"2025-12-29T20:03:53.341187Z","shell.execute_reply":"2025-12-29T20:03:53.358408Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef evaluate_model(model, loader):\n    model.eval()\n    all_preds = []\n    all_labels = []\n\n    pbar = tqdm(loader, leave=False)\n    pbar.set_description('Evaluating')\n\n    for images, labels in pbar:\n        images = images.to(device, non_blocking=True)\n        labels = labels.to(device, non_blocking=True)\n\n        outputs = model(images)\n        preds = outputs.argmax(dim=1)\n\n        all_preds.append(preds.cpu().numpy())\n        all_labels.append(labels.cpu().numpy())\n\n    preds = np.concatenate(all_preds)\n    labels = np.concatenate(all_labels)\n\n    accuracy, f1, precision, recall, cmat = compute_classification_metrics(preds, labels)\n    return preds, labels, accuracy, f1, precision, recall, cmat\n\n\n@torch.no_grad()\ndef generate_submission(model, loader):\n    model.eval()\n    all_preds = []\n    all_ids = []\n\n    pbar = tqdm(loader, leave=False)\n    pbar.set_description('Predicting test')\n\n    for images, image_ids in pbar:\n        images = images.to(device, non_blocking=True)\n        outputs = model(images)\n        preds = outputs.argmax(dim=1).cpu().numpy()\n        all_preds.extend(preds.tolist())\n        all_ids.extend(image_ids)\n\n    submission = pd.DataFrame({'id': all_ids, 'label': all_preds})\n    submission.to_csv(ModelConfig.submission_path, index=False)\n    print(f'Submission saved to {ModelConfig.submission_path}')\n    return submission\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-29T20:03:53.360012Z","iopub.execute_input":"2025-12-29T20:03:53.360377Z","iopub.status.idle":"2025-12-29T20:03:53.382076Z","shell.execute_reply.started":"2025-12-29T20:03:53.360336Z","shell.execute_reply":"2025-12-29T20:03:53.381390Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\n\nmodel, history, best_val_acc = train_model()\nplot_training_curves(history)\nprint(f'Best validation accuracy: {best_val_acc:.4f}')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-29T20:03:53.384157Z","iopub.execute_input":"2025-12-29T20:03:53.384350Z","iopub.status.idle":"2025-12-29T22:34:02.981919Z","shell.execute_reply.started":"2025-12-29T20:03:53.384331Z","shell.execute_reply":"2025-12-29T22:34:02.981058Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"checkpoint = torch.load(ModelConfig.model_path, map_location=device)\nbase_model = build_model()\nbase_model.load_state_dict(checkpoint['model_state_dict'])\nif use_cuda and num_gpus > 1:\n    model = nn.DataParallel(base_model).to(device)\nelse:\n    model = base_model.to(device)\nprint(f\"Loaded best checkpoint with val_acc={checkpoint['val_acc']:.4f}\")\n\n_, val_loader = get_loaders()\npreds, labels, acc, f1, precision, recall, cmat = evaluate_model(model, val_loader)\nprint(f'Validation accuracy: {acc:.4f}, F1: {f1:.4f}, precision: {precision:.4f}, recall: {recall:.4f}')\n\ncmat_norm = (cmat.T / cmat.sum(axis=1, keepdims=True)).T\nshow_confusion_matrix(cmat_norm, title_suffix='(normalized)')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-29T22:34:02.983378Z","iopub.execute_input":"2025-12-29T22:34:02.983675Z","iopub.status.idle":"2025-12-29T22:37:17.853021Z","shell.execute_reply.started":"2025-12-29T22:34:02.983642Z","shell.execute_reply":"2025-12-29T22:37:17.852128Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_loader = get_test_loader()\nsubmission = generate_submission(model, test_loader)\nsubmission.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-29T22:37:17.854858Z","iopub.execute_input":"2025-12-29T22:37:17.855222Z","iopub.status.idle":"2025-12-29T22:43:22.392128Z","shell.execute_reply.started":"2025-12-29T22:37:17.855189Z","shell.execute_reply":"2025-12-29T22:43:22.391392Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}