{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"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":"none","dataSources":[{"sourceType":"competition","sourceId":21154,"databundleVersionId":1243559}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 🌸 Visual Flower Classification Pipeline — TPU Getting Started / Oxford Flowers 104\n\nThis notebook is a fully rewritten, visual-first version of the original training script.\n\nIt focuses on:\n\n- Clear step-by-step cells\n- Rich visual explanations with `seaborn` and `matplotlib`\n- Safer Kaggle path handling without `kagglehub.login()`\n- TFRecord reading for JPEG images\n- PyTorch + `timm` image classification\n- Training curves, confusion matrix, class diagnostics, and sample predictions\n\n> Designed for Kaggle Notebook execution. For a quick smoke test, set `DEBUG = True` in the configuration cell.","metadata":{}},{"cell_type":"markdown","source":"## 1. Install and import libraries\n\nKaggle usually has many packages preinstalled, but `tfrecord` and sometimes `timm` may need installation.","metadata":{}},{"cell_type":"code","source":"# CELL 1: Install dependencies\n!pip install -q tfrecord timm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T13:06:11.310347Z","iopub.execute_input":"2026-04-29T13:06:11.311174Z","iopub.status.idle":"2026-04-29T13:06:16.021472Z","shell.execute_reply.started":"2026-04-29T13:06:11.311128Z","shell.execute_reply":"2026-04-29T13:06:16.020542Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CELL 2: Imports\n# CELL 2: Imports\nimport warnings\n\nwarnings.filterwarnings(\"ignore\", category=UserWarning)\nwarnings.filterwarnings(\"ignore\", message=\".*UnsupportedFieldAttributeWarning.*\")\nwarnings.filterwarnings(\"ignore\", module=\"pydantic.*\")\n\nimport io\nimport math\nimport os\nimport random\nimport re\nimport warnings\nfrom pathlib import Path\nfrom collections import Counter\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\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 (\n    accuracy_score,\n    f1_score,\n    precision_score,\n    recall_score,\n    confusion_matrix,\n    classification_report,\n)\n\nimport timm\nimport tfrecord\nfrom tqdm.auto import tqdm\n\nwarnings.filterwarnings('ignore')\nsns.set_theme(style='whitegrid', context='notebook')\nplt.rcParams['figure.figsize'] = (12, 6)\nplt.rcParams['axes.titleweight'] = 'bold'\n\nprint('Torch:', torch.__version__)\nprint('CUDA available:', torch.cuda.is_available())\nprint('Device count:', torch.cuda.device_count())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T13:06:16.023152Z","iopub.execute_input":"2026-04-29T13:06:16.023398Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Configuration\n\nThe original code used a large Swin model. That is strong but heavy. This version keeps it configurable:\n\n- `DEBUG = True`: fast visual/debug run\n- `DEBUG = False`: full training run\n- `MODEL_NAME`: can be changed to a larger model if GPU time is sufficient\n\nRecommended starting point on Kaggle: run debug first, then full training.","metadata":{}},{"cell_type":"code","source":"# CELL 3: Configuration\nclass CFG:\n    # Kaggle competition input path\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    OUTPUT_DIR = Path('/kaggle/working')\n    OUTPUT_DIR.mkdir(parents=True, exist_ok=True)\n    MODEL_PATH = OUTPUT_DIR / 'best_flower_model.pt'\n    SUBMISSION_PATH = OUTPUT_DIR / 'submission.csv'\n\n    # Switch this off for the full run\n    DEBUG = False\n\n    IMAGE_SIZE = 384\n    NUM_CLASSES = 104\n\n    # Large model from the original code. If runtime is too slow, use:\n    # 'tf_efficientnet_b3_ns' or 'convnext_tiny.fb_in22k_ft_in1k'\n    MODEL_NAME = 'swin_large_patch4_window12_384'\n\n    BATCH_SIZE = 16\n    VAL_BATCH_SIZE = 64\n    TEST_BATCH_SIZE = 64\n    EPOCHS = 10\n\n    LR = 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    USE_AMP = True\n\n    DROP_RATE = 0.214\n    DROP_PATH_RATE = 0.118\n\n    # Debug limits\n    TRAIN_LIMIT = 512 if DEBUG else None\n    VAL_LIMIT = 256 if DEBUG else None\n    TEST_LIMIT = 256 if DEBUG else None\n\nCLASSES = [\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\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nUSE_CUDA = torch.cuda.is_available()\nNUM_GPUS = torch.cuda.device_count() if USE_CUDA else 0\n\nprint('Device:', DEVICE)\nprint('Debug mode:', CFG.DEBUG)\nprint('Model:', CFG.MODEL_NAME)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Reproducibility and path check\n\nThis cell checks that the competition data is attached to the Kaggle Notebook.","metadata":{}},{"cell_type":"code","source":"# CELL 4: Reproducibility and path validation\ndef set_seed(seed: int = CFG.SEED):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False\n\nset_seed()\n\nrequired_dirs = [CFG.TRAIN_DIR, CFG.VAL_DIR, CFG.TEST_DIR]\nfor d in required_dirs:\n    print(f'{str(d):70s} exists={d.exists()}')\n\nif not CFG.BASE_PATH.exists():\n    raise FileNotFoundError(\n        'Competition data not found. On Kaggle, add the competition dataset: tpu-getting-started.'\n    )","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. TFRecord file overview\n\nThe item count is encoded in the TFRecord filename, so we can estimate dataset size without loading all records.","metadata":{}},{"cell_type":"code","source":"# CELL 5: TFRecord utilities and dataset size overview\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\ntrain_files = get_tfrecord_files(CFG.TRAIN_DIR)\nval_files = get_tfrecord_files(CFG.VAL_DIR)\ntest_files = get_tfrecord_files(CFG.TEST_DIR)\n\nsize_df = pd.DataFrame({\n    'split': ['train', 'validation', 'test'],\n    'num_tfrecord_files': [len(train_files), len(val_files), len(test_files)],\n    'estimated_items': [count_data_items(train_files), count_data_items(val_files), count_data_items(test_files)],\n})\nsize_df","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CELL 6: Visualize split sizes\nfig, axes = plt.subplots(1, 2, figsize=(15, 5))\n\nsns.barplot(data=size_df, x='split', y='num_tfrecord_files', ax=axes[0])\naxes[0].set_title('Number of TFRecord Files by Split')\naxes[0].set_xlabel('Split')\naxes[0].set_ylabel('Files')\n\nsns.barplot(data=size_df, x='split', y='estimated_items', ax=axes[1])\naxes[1].set_title('Estimated Number of Images by Split')\naxes[1].set_xlabel('Split')\naxes[1].set_ylabel('Images')\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Read and visualize raw examples\n\nBefore training, inspect actual images and labels. This makes it easier to catch data reading or label problems early.","metadata":{}},{"cell_type":"code","source":"# CELL 7: Raw TFRecord reader for visualization\ndef parse_label(value):\n    if isinstance(value, np.ndarray):\n        return int(value.reshape(-1)[0])\n    if isinstance(value, (list, tuple)):\n        return int(value[0])\n    return int(value)\n\n\ndef read_raw_samples(filenames, n=16, labeled=True):\n    samples = []\n    for fname in filenames:\n        for record in tfrecord.tfrecord_loader(fname, None, None):\n            image = Image.open(io.BytesIO(record['image'])).convert('RGB')\n            if labeled:\n                label = parse_label(record['class'])\n                samples.append((image, label))\n            else:\n                image_id = record.get('id', '')\n                if isinstance(image_id, bytes):\n                    image_id = image_id.decode()\n                samples.append((image, image_id))\n            if len(samples) >= n:\n                return samples\n    return samples\n\nraw_samples = read_raw_samples(train_files, n=20, labeled=True)\nprint('Loaded raw samples:', len(raw_samples))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CELL 8: Display raw image grid\ndef show_image_grid(samples, class_names=None, ncols=5, title='Raw Training Samples'):\n    n = len(samples)\n    nrows = math.ceil(n / ncols)\n    fig, axes = plt.subplots(nrows, ncols, figsize=(3.2 * ncols, 3.4 * nrows))\n    axes = np.array(axes).reshape(-1)\n\n    for ax, item in zip(axes, samples):\n        image, label = item\n        ax.imshow(image)\n        if class_names is not None and isinstance(label, int):\n            label_text = f'{label}: {class_names[label]}'\n        else:\n            label_text = str(label)\n        ax.set_title(label_text, fontsize=9)\n        ax.axis('off')\n\n    for ax in axes[len(samples):]:\n        ax.axis('off')\n\n    fig.suptitle(title, fontsize=18, fontweight='bold')\n    plt.tight_layout()\n    plt.show()\n\nshow_image_grid(raw_samples, CLASSES, ncols=5)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Label distribution preview\n\nThis is a lightweight scan. For a complete count, increase `max_records` or scan the full training set.","metadata":{}},{"cell_type":"code","source":"# CELL 9: Label distribution scan\nCOLUMNS_TO_SHOW = 25\n\ndef scan_label_distribution(filenames, max_records=5000):\n    counter = Counter()\n    seen = 0\n    for fname in tqdm(filenames, desc='Scanning labels'):\n        for record in tfrecord.tfrecord_loader(fname, None, None):\n            label = parse_label(record['class'])\n            counter[label] += 1\n            seen += 1\n            if seen >= max_records:\n                break\n        if seen >= max_records:\n            break\n    df = pd.DataFrame({\n        'label': list(counter.keys()),\n        'count': list(counter.values()),\n    }).sort_values('count', ascending=False)\n    df['class_name'] = df['label'].map(lambda x: CLASSES[x])\n    return df\n\nlabel_df = scan_label_distribution(train_files, max_records=5000 if not CFG.DEBUG else 1000)\nlabel_df.head(10)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CELL 10: Visualize top class frequencies\nplt.figure(figsize=(14, 7))\nsns.barplot(data=label_df.head(COLUMNS_TO_SHOW), y='class_name', x='count')\nplt.title(f'Top {COLUMNS_TO_SHOW} Classes in Scanned Training Records')\nplt.xlabel('Count in scanned records')\nplt.ylabel('Class')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Augmentation preview\n\nThe training transforms intentionally create variation: crop, flip, rotation, color jitter, normalization, and random erasing.","metadata":{}},{"cell_type":"code","source":"# CELL 11: Transforms\ndef get_train_transforms():\n    return transforms.Compose([\n        transforms.RandomResizedCrop((CFG.IMAGE_SIZE, CFG.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((CFG.IMAGE_SIZE, CFG.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\ndef denormalize_tensor(t):\n    mean = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1)\n    std = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1)\n    x = t.cpu() * std + mean\n    return x.clamp(0, 1).permute(1, 2, 0).numpy()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CELL 12: Augmentation visual preview\nsample_img, sample_label = raw_samples[0]\ntrain_tfms = get_train_transforms()\n\naugmented = [(train_tfms(sample_img), sample_label) for _ in range(12)]\nfig, axes = plt.subplots(3, 4, figsize=(14, 10))\nfor ax, (tensor_img, label) in zip(axes.ravel(), augmented):\n    ax.imshow(denormalize_tensor(tensor_img))\n    ax.set_title(CLASSES[label], fontsize=9)\n    ax.axis('off')\nplt.suptitle('Random Augmentation Preview', fontsize=18, fontweight='bold')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8. Dataset and DataLoader\n\nThis uses an `IterableDataset` because TFRecord files are streamed record by record.","metadata":{}},{"cell_type":"code","source":"# CELL 13: Dataset and loaders\nclass FlowerTFRecordDataset(IterableDataset):\n    def __init__(self, filenames, labeled=True, transform=None, seed=CFG.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        # Stable hash alternative to avoid Python hash randomization across sessions\n        stable = sum(ord(c) for c in str(fname)) % 10_000\n        rng = random.Random(self.seed + self.epoch + stable)\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 = Image.open(io.BytesIO(record['image'])).convert('RGB')\n                if self.transform is not None:\n                    image = self.transform(image)\n\n                if self.labeled:\n                    label = parse_label(record['class'])\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    loader = DataLoader(\n        dataset,\n        batch_size=batch_size,\n        shuffle=False,\n        num_workers=CFG.NUM_WORKERS,\n        pin_memory=torch.cuda.is_available(),\n        drop_last=False,\n    )\n    return loader, dataset\n\n\ndef get_loaders():\n    train_loader, train_dataset = build_loader(\n        CFG.TRAIN_DIR,\n        labeled=True,\n        transform=get_train_transforms(),\n        batch_size=CFG.BATCH_SIZE,\n        limit=CFG.TRAIN_LIMIT,\n    )\n    val_loader, val_dataset = build_loader(\n        CFG.VAL_DIR,\n        labeled=True,\n        transform=get_eval_transforms(),\n        batch_size=CFG.VAL_BATCH_SIZE,\n        limit=CFG.VAL_LIMIT,\n    )\n    return train_loader, val_loader, train_dataset, val_dataset\n\n\ndef get_test_loader():\n    test_loader, test_dataset = build_loader(\n        CFG.TEST_DIR,\n        labeled=False,\n        transform=get_eval_transforms(),\n        batch_size=CFG.TEST_BATCH_SIZE,\n        limit=CFG.TEST_LIMIT,\n    )\n    return test_loader, test_dataset\n\ntrain_loader, val_loader, train_dataset, val_dataset = get_loaders()\nprint('Train samples:', len(train_dataset))\nprint('Validation samples:', len(val_dataset))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 9. Model, optimizer, scheduler, and helper metrics","metadata":{}},{"cell_type":"code","source":"# CELL 14: Model and training helpers\nclass AverageMeter:\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 build_model():\n    model = timm.create_model(\n        CFG.MODEL_NAME,\n        pretrained=True,\n        num_classes=CFG.NUM_CLASSES,\n        drop_rate=CFG.DROP_RATE,\n        drop_path_rate=CFG.DROP_PATH_RATE,\n    )\n    return model\n\n\ndef create_optimizer(model):\n    return torch.optim.AdamW(\n        model.parameters(),\n        lr=CFG.LR,\n        weight_decay=CFG.WEIGHT_DECAY,\n    )\n\n\ndef create_scheduler(optimizer):\n    return torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer,\n        T_max=CFG.EPOCHS,\n        eta_min=CFG.MIN_LR,\n    )","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 10. Training and validation loop","metadata":{}},{"cell_type":"code","source":"# CELL 15: One epoch runner\ndef 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    model.train() if training else model.eval()\n\n    losses = AverageMeter()\n    accuracies = AverageMeter()\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\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                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.GRAD_CLIP)\n                scaler.step(optimizer)\n                scaler.update()\n            else:\n                loss.backward()\n                torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.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        lr_value = optimizer.param_groups[0]['lr'] if optimizer is not None else 0.0\n        pbar.set_postfix(loss=f'{losses.avg:.4f}', acc=f'{accuracies.avg:.4f}', lr=f'{lr_value:.2e}')\n\n    return losses.avg, accuracies.avg","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CELL 16: Full training function\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\n    criterion = nn.CrossEntropyLoss()\n    optimizer = create_optimizer(model)\n    scheduler = create_scheduler(optimizer)\n    scaler = GradScaler(enabled=CFG.USE_AMP and DEVICE.type == 'cuda')\n\n    history = {\n        'epoch': [],\n        'train_loss': [],\n        'train_acc': [],\n        'val_loss': [],\n        'val_acc': [],\n        'lr': [],\n    }\n\n    best_val_acc = -1.0\n\n    for epoch in range(CFG.EPOCHS):\n        train_loss, train_acc = run_epoch(\n            train_loader, model, criterion,\n            optimizer=optimizer, scaler=scaler,\n            epoch=epoch, training=True,\n        )\n\n        val_loss, val_acc = run_epoch(\n            val_loader, model, criterion,\n            optimizer=None, scaler=None,\n            epoch=epoch, training=False,\n        )\n\n        scheduler.step()\n        current_lr = optimizer.param_groups[0]['lr']\n\n        history['epoch'].append(epoch + 1)\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(\n            f\"Epoch {epoch + 1}/{CFG.EPOCHS} | \"\n            f\"train_loss={train_loss:.4f}, train_acc={train_acc:.4f} | \"\n            f\"val_loss={val_loss:.4f}, val_acc={val_acc:.4f} | lr={current_lr:.2e}\"\n        )\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({\n                'model_state_dict': model_to_save.state_dict(),\n                'val_acc': best_val_acc,\n                'cfg': {k: v for k, v in CFG.__dict__.items() if not k.startswith('__')},\n            }, CFG.MODEL_PATH)\n            print(f'  ✓ New best model saved: {CFG.MODEL_PATH} | val_acc={best_val_acc:.4f}')\n\n    history_df = pd.DataFrame(history)\n    history_df.to_csv(CFG.OUTPUT_DIR / 'training_history.csv', index=False)\n    return model, history_df, best_val_acc","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 11. Visual training diagnostics","metadata":{}},{"cell_type":"code","source":"# CELL 17: Training visualization functions\ndef plot_training_dashboard(history_df):\n    fig, axes = plt.subplots(1, 3, figsize=(20, 5))\n\n    sns.lineplot(data=history_df, x='epoch', y='train_loss', marker='o', label='train', ax=axes[0])\n    sns.lineplot(data=history_df, x='epoch', y='val_loss', marker='o', label='validation', ax=axes[0])\n    axes[0].set_title('Loss Curve')\n    axes[0].set_xlabel('Epoch')\n    axes[0].set_ylabel('Loss')\n\n    sns.lineplot(data=history_df, x='epoch', y='train_acc', marker='o', label='train', ax=axes[1])\n    sns.lineplot(data=history_df, x='epoch', y='val_acc', marker='o', label='validation', ax=axes[1])\n    axes[1].set_title('Accuracy Curve')\n    axes[1].set_xlabel('Epoch')\n    axes[1].set_ylabel('Accuracy')\n\n    sns.lineplot(data=history_df, x='epoch', y='lr', marker='o', ax=axes[2])\n    axes[2].set_title('Learning Rate Schedule')\n    axes[2].set_xlabel('Epoch')\n    axes[2].set_ylabel('Learning Rate')\n    axes[2].set_yscale('log')\n\n    plt.tight_layout()\n    plt.show()\n\n\ndef show_best_epoch_card(history_df):\n    best_row = history_df.loc[history_df['val_acc'].idxmax()]\n    summary = pd.DataFrame({\n        'metric': ['Best epoch', 'Best validation accuracy', 'Validation loss at best epoch', 'Training accuracy at best epoch'],\n        'value': [\n            int(best_row['epoch']),\n            f\"{best_row['val_acc']:.4f}\",\n            f\"{best_row['val_loss']:.4f}\",\n            f\"{best_row['train_acc']:.4f}\",\n        ]\n    })\n    display(summary)\n    return summary","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 12. Train the model\n\nRun this cell to train. It saves the best checkpoint to `/kaggle/working/best_flower_model.pt`.","metadata":{}},{"cell_type":"code","source":"# CELL 18: Train\n\nmodel, history_df, best_val_acc = train_model()\nplot_training_dashboard(history_df)\nshow_best_epoch_card(history_df)\nprint(f'Best validation accuracy: {best_val_acc:.4f}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 13. Load the best checkpoint and evaluate validation performance","metadata":{}},{"cell_type":"code","source":"# CELL 19: Load best checkpoint\ndef load_best_model():\n    checkpoint = torch.load(CFG.MODEL_PATH, map_location=DEVICE)\n    base_model = build_model()\n    base_model.load_state_dict(checkpoint['model_state_dict'])\n    if USE_CUDA and NUM_GPUS > 1:\n        model = nn.DataParallel(base_model).to(DEVICE)\n    else:\n        model = base_model.to(DEVICE)\n    print(f\"Loaded best checkpoint: val_acc={checkpoint['val_acc']:.4f}\")\n    return model\n\nbest_model = load_best_model()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CELL 20: Validation evaluation\n@torch.no_grad()\ndef evaluate_model(model, loader):\n    model.eval()\n    all_preds = []\n    all_labels = []\n    all_probs = []\n\n    pbar = tqdm(loader, leave=False, desc='Evaluating validation')\n    for images, labels in pbar:\n        images = images.to(DEVICE, non_blocking=True)\n        labels = labels.to(DEVICE, non_blocking=True)\n\n        logits = model(images)\n        probs = torch.softmax(logits, dim=1)\n        preds = probs.argmax(dim=1)\n\n        all_preds.append(preds.cpu().numpy())\n        all_labels.append(labels.cpu().numpy())\n        all_probs.append(probs.cpu().numpy())\n\n    preds = np.concatenate(all_preds)\n    labels = np.concatenate(all_labels)\n    probs = np.concatenate(all_probs)\n\n    metrics = {\n        'accuracy': accuracy_score(labels, preds),\n        'macro_f1': f1_score(labels, preds, labels=np.arange(CFG.NUM_CLASSES), average='macro', zero_division=0),\n        'macro_precision': precision_score(labels, preds, labels=np.arange(CFG.NUM_CLASSES), average='macro', zero_division=0),\n        'macro_recall': recall_score(labels, preds, labels=np.arange(CFG.NUM_CLASSES), average='macro', zero_division=0),\n    }\n    cmat = confusion_matrix(labels, preds, labels=np.arange(CFG.NUM_CLASSES))\n    return preds, labels, probs, metrics, cmat\n\n_, val_loader, _, _ = get_loaders()\npreds, labels, probs, metrics, cmat = evaluate_model(best_model, val_loader)\nmetrics_df = pd.DataFrame(metrics.items(), columns=['metric', 'value'])\nmetrics_df","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CELL 21: Metrics bar chart\nplt.figure(figsize=(10, 5))\nsns.barplot(data=metrics_df, x='metric', y='value')\nplt.ylim(0, 1)\nplt.title('Validation Metrics')\nplt.xlabel('Metric')\nplt.ylabel('Score')\nplt.xticks(rotation=20)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 14. Confusion matrix and weakest classes\n\nA 104×104 matrix is large, so this notebook shows both the full normalized matrix and a focused view of the weakest classes.","metadata":{}},{"cell_type":"code","source":"# CELL 22: Confusion matrix visualizations\ndef plot_confusion_matrix(cmat, normalize=True, max_classes=None, title='Confusion Matrix'):\n    matrix = cmat.astype(float)\n    if normalize:\n        denom = matrix.sum(axis=1, keepdims=True)\n        matrix = np.divide(matrix, denom, out=np.zeros_like(matrix), where=denom != 0)\n\n    if max_classes is not None:\n        matrix = matrix[:max_classes, :max_classes]\n        labels_short = [str(i) for i in range(max_classes)]\n    else:\n        labels_short = [str(i) for i in range(matrix.shape[0])]\n\n    plt.figure(figsize=(16, 14))\n    sns.heatmap(matrix, cmap='Reds', xticklabels=labels_short, yticklabels=labels_short, square=True)\n    plt.title(title)\n    plt.xlabel('Predicted label')\n    plt.ylabel('True label')\n    plt.tight_layout()\n    plt.show()\n\nplot_confusion_matrix(cmat, normalize=True, max_classes=None, title='Normalized Confusion Matrix — All Classes')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CELL 23: Per-class accuracy and weakest classes\nclass_totals = cmat.sum(axis=1)\nclass_correct = np.diag(cmat)\nclass_acc = np.divide(class_correct, class_totals, out=np.zeros_like(class_correct, dtype=float), where=class_totals != 0)\n\nclass_perf_df = pd.DataFrame({\n    'label': np.arange(CFG.NUM_CLASSES),\n    'class_name': CLASSES,\n    'support': class_totals,\n    'correct': class_correct,\n    'class_accuracy': class_acc,\n}).sort_values('class_accuracy')\n\nweak_df = class_perf_df[class_perf_df['support'] > 0].head(20)\ndisplay(weak_df)\n\nplt.figure(figsize=(14, 8))\nsns.barplot(data=weak_df, y='class_name', x='class_accuracy')\nplt.title('20 Weakest Classes by Validation Accuracy')\nplt.xlabel('Class Accuracy')\nplt.ylabel('Class')\nplt.xlim(0, 1)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CELL 24: Most common mistakes\nmistakes = []\nfor true_label in range(CFG.NUM_CLASSES):\n    for pred_label in range(CFG.NUM_CLASSES):\n        if true_label != pred_label and cmat[true_label, pred_label] > 0:\n            mistakes.append({\n                'true_label': true_label,\n                'true_class': CLASSES[true_label],\n                'pred_label': pred_label,\n                'pred_class': CLASSES[pred_label],\n                'count': int(cmat[true_label, pred_label]),\n            })\n\nmistake_df = pd.DataFrame(mistakes).sort_values('count', ascending=False).head(25)\ndisplay(mistake_df)\n\nif len(mistake_df) > 0:\n    mistake_df['pair'] = mistake_df['true_class'] + ' → ' + mistake_df['pred_class']\n    plt.figure(figsize=(14, 9))\n    sns.barplot(data=mistake_df, y='pair', x='count')\n    plt.title('Most Frequent Validation Mistakes')\n    plt.xlabel('Mistake Count')\n    plt.ylabel('True → Predicted')\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 15. Submission generation","metadata":{}},{"cell_type":"code","source":"# CELL 25: Generate submission\n@torch.no_grad()\ndef generate_submission(model):\n    test_loader, _ = get_test_loader()\n    model.eval()\n    all_preds = []\n    all_ids = []\n\n    pbar = tqdm(test_loader, leave=False, desc='Predicting test')\n    for images, image_ids in pbar:\n        images = images.to(DEVICE, non_blocking=True)\n        logits = model(images)\n        preds = logits.argmax(dim=1).cpu().numpy()\n        all_preds.extend(preds.tolist())\n        all_ids.extend(list(image_ids))\n\n    submission = pd.DataFrame({'id': all_ids, 'label': all_preds})\n    submission.to_csv(CFG.SUBMISSION_PATH, index=False)\n    print(f'Submission saved to: {CFG.SUBMISSION_PATH}')\n    return submission\n\nsubmission = generate_submission(best_model)\nsubmission.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CELL 26: Submission prediction distribution\npred_dist = submission['label'].value_counts().reset_index()\npred_dist.columns = ['label', 'count']\npred_dist['class_name'] = pred_dist['label'].map(lambda x: CLASSES[int(x)] if int(x) < len(CLASSES) else str(x))\n\ndisplay(pred_dist.head(20))\n\nplt.figure(figsize=(14, 7))\nsns.barplot(data=pred_dist.head(25), y='class_name', x='count')\nplt.title('Top Predicted Classes in Test Submission')\nplt.xlabel('Prediction count')\nplt.ylabel('Class')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 16. Final checklist\n\nAfter running the notebook, confirm:\n\n1. `best_flower_model.pt` exists in `/kaggle/working`\n2. `training_history.csv` exists in `/kaggle/working`\n3. `submission.csv` exists in `/kaggle/working`\n4. The submission file has two columns: `id`, `label`\n5. The label values are integers from 0 to 103","metadata":{}},{"cell_type":"code","source":"# CELL 27: Final file check\nfor path in [CFG.MODEL_PATH, CFG.OUTPUT_DIR / 'training_history.csv', CFG.SUBMISSION_PATH]:\n    print(f'{path}: exists={path.exists()}')\n\nif CFG.SUBMISSION_PATH.exists():\n    sub_check = pd.read_csv(CFG.SUBMISSION_PATH)\n    display(sub_check.head())\n    print('Shape:', sub_check.shape)\n    print('Columns:', list(sub_check.columns))\n    print('Label range:', sub_check['label'].min(), 'to', sub_check['label'].max())","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}