{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":21154,"databundleVersionId":1243559}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Flower Classification\n\nThis notebook explores flower classification on the [Petals to the Metal](https://www.kaggle.com/c/tpu-getting-started) Kaggle competition, using a Vision Transformer (ViT).\n\nHere is a summary of the lessons learned:\n* Ensemble heads don't help; a full-model ensemble might, but is likely too resource-intensive.\n* Sequentially unfreezing layers during fine-tuning improved performance.\n* A cosine decay learning rate schedule with warm-up yielded better fine-tuning results.\n* Data augmentation helped on the original dataset but appeared to confuse the model on extended data.\n* This version uses **PyTorch** instead of TensorFlow/Keras.","metadata":{}},{"cell_type":"markdown","source":"# Setting up the notebook","metadata":{}},{"cell_type":"code","source":"!pip install transformers torch torchvision tfrecord scikit-learn","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"For a P100 we need CUDA 11.8. The P100 is a Pascal architecture GPU (compute capability 6.0) and doesn't support CUDA 12.x kernels. The default PyTorch build on this Kaggle environment only supports sm_70+ (Volta and newer), and the P100 is sm_60 (Pascal). No available pre-built PyTorch wheel supports sm_60 anymore in recent versions.\n","metadata":{}},{"cell_type":"code","source":"# !pip install torch==1.13.1+cu116 torchvision==0.14.1+cu116 --index-url https://download.pytorch.org/whl/cu116 --upgrade","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport re\nimport io\nimport math\nimport matplotlib.pyplot as plt\nfrom collections import Counter\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader, Dataset\nimport torchvision.transforms as T\nfrom PIL import Image\n\nfrom transformers import AutoModel\n\nprint(\"PyTorch version:\", torch.__version__)\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", DEVICE)\ntorch.manual_seed(42)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## A note on ml_utils 🛠️\n\nThroughout this notebook, I use [ml_utils](https://github.com/ThomasPRZilliox/ml_utils): a small personal toolkit \nI've been building to keep ML notebooks clean. The PyTorch utilities are a work in progress — visualisation helpers are implemented inline below.\n\nIf you find it useful, a ⭐ on GitHub is always appreciated! And if something doesn't work, feel free to [open an issue](https://github.com/ThomasPRZilliox/ml_utils/issues).","metadata":{}},{"cell_type":"markdown","source":"## Device Detection\nPyTorch uses CUDA for GPU acceleration. Multi-GPU training can be enabled via `DataParallel` or `DistributedDataParallel`.","metadata":{}},{"cell_type":"code","source":"print(f\"GPUs available: {torch.cuda.device_count()}\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Prepare the dataset","metadata":{}},{"cell_type":"markdown","source":"## Get the dataset from Kaggle","metadata":{}},{"cell_type":"code","source":"IMAGE_SIZE = [224, 224]\nGCS_PATH = '/kaggle/input/competitions/tpu-getting-started/tfrecords-jpeg-224x224'\n\ndef list_tfrec(path):\n    if os.path.isdir(path):\n        return sorted([os.path.join(path, f) for f in os.listdir(path) if f.endswith('.tfrec')])\n    return []\n\nTRAINING_FILENAMES   = list_tfrec(GCS_PATH + '/train')\nVALIDATION_FILENAMES = list_tfrec(GCS_PATH + '/val')\nTEST_FILENAMES       = list_tfrec(GCS_PATH + '/test')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Extra training data\n\nTo boost the training set, we also use the external dataset [tf-flower-photo-tfrec](https://www.kaggle.com/datasets/kirillblinov/tf-flower-photo-tfrec). The `_no_test` folders ensure none of the competition test images leak into training.","metadata":{}},{"cell_type":"code","source":"GCS_PATH_imagenet    = '/kaggle/input/datasets/kirillblinov/tf-flower-photo-tfrec/imagenet/tfrecords-jpeg-224x224'\nGCS_PATH_inaturalist = '/kaggle/input/datasets/kirillblinov/tf-flower-photo-tfrec/inaturalist/tfrecords-jpeg-224x224'\nGCS_PATH_openimage   = '/kaggle/input/datasets/kirillblinov/tf-flower-photo-tfrec/openimage/tfrecords-jpeg-224x224'\nGCS_PATH_oxford_102  = '/kaggle/input/datasets/kirillblinov/tf-flower-photo-tfrec/oxford_102/tfrecords-jpeg-224x224'\n\nIMAGENET_FILENAMES    = list_tfrec(GCS_PATH_imagenet)\nINATURALIST_FILENAMES = list_tfrec(GCS_PATH_inaturalist)\nOPENIMAGE_FILENAMES   = list_tfrec(GCS_PATH_openimage)\nOXFORD_102_FILENAMES  = list_tfrec(GCS_PATH_oxford_102)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TRAINING_FILENAMES = TRAINING_FILENAMES + IMAGENET_FILENAMES + INATURALIST_FILENAMES + OPENIMAGE_FILENAMES + OXFORD_102_FILENAMES","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Define the classes","metadata":{}},{"cell_type":"code","source":"CLASSES = ['pink primrose',    'hard-leaved pocket orchid', 'canterbury bells', 'sweet pea',     'wild geranium',     'tiger lily',           'moon orchid',              'bird of paradise', 'monkshood',        'globe thistle',         # 00 - 09\n           'snapdragon',       \"colt's foot\",               'king protea',      'spear thistle', 'yellow iris',       'globe-flower',         'purple coneflower',        'peruvian lily',    'balloon flower',   'giant white arum lily', # 10 - 19\n           'fire lily',        'pincushion flower',         'fritillary',       'red ginger',    'grape hyacinth',    'corn poppy',           'prince of wales feathers', 'stemless gentian', 'artichoke',        'sweet william',         # 20 - 29\n           'carnation',        'garden phlox',              'love in the mist', 'cosmos',        'alpine sea holly',  'ruby-lipped cattleya', 'cape flower',              'great masterwort', 'siam tulip',       'lenten rose',           # 30 - 39\n           'barberton daisy',  'daffodil',                  'sword lily',       'poinsettia',    'bolero deep blue',  'wallflower',           'marigold',                 'buttercup',        'daisy',            'common dandelion',      # 40 - 49\n           'petunia',          'wild pansy',                'primula',          'sunflower',     'lilac hibiscus',    'bishop of llandaff',   'gaura',                    'geranium',         'orange dahlia',    'pink-yellow dahlia',    # 50 - 59\n           'cautleya spicata', 'japanese anemone',          'black-eyed susan', 'silverbush',    'californian poppy', 'osteospermum',         'spring crocus',            'iris',             'windflower',       'tree poppy',            # 60 - 69\n           'gazania',          'azalea',                    'water lily',       'rose',          'thorn apple',       'morning glory',        'passion flower',           'lotus',            'toad lily',        'anthurium',             # 70 - 79\n           'frangipani',       'clematis',                  'hibiscus',         'columbine',     'desert-rose',       'tree mallow',          'magnolia',                 'cyclamen ',        'watercress',       'canna lily',            # 80 - 89\n           'hippeastrum ',     'bee balm',                  'pink quill',       'foxglove',      'bougainvillea',     'camellia',             'mallow',                   'mexican petunia',  'bromelia',         'blanket flower',        # 90 - 99\n           'trumpet creeper',  'blackberry lily',           'common tulip',     'wild rose']                                                                                                                                               # 100 - 102","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Definition of dataset helper functions\n\nPyTorch has no native TFRecord reader, so we use the `tfrecord` package to parse `.tfrec` files. Images are decoded with PIL and normalized using standard ImageNet mean/std, identical to the original notebook.","metadata":{}},{"cell_type":"code","source":"from tfrecord.torch.dataset import TFRecordDataset\n\nIMAGENET_MEAN = [0.485, 0.456, 0.406]\nIMAGENET_STD  = [0.229, 0.224, 0.225]\n\n\ndef count_data_items(filenames):\n    \"\"\"Infer sample count from filename, e.g. flowers00-230.tfrec -> 230.\"\"\"\n    n = [int(re.compile(r'-(\\d+)\\.').search(f).group(1))\n         for f in filenames if re.search(r'-(\\d+)\\.', f)]\n    return int(np.sum(n))\n\n\nclass FlowerDataset(Dataset):\n    \"\"\"Wraps one or more TFRecord files as a PyTorch Dataset.\"\"\"\n\n    def __init__(self, filenames, labeled=True, augment=False):\n        self.labeled = labeled\n        self.samples = []   # list of (jpeg_bytes, label_int | id_str)\n\n        base_tfm = T.Compose([\n            T.Resize(IMAGE_SIZE),\n            T.ToTensor(),\n            T.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n        ])\n        aug_tfm = T.Compose([\n            T.Resize(IMAGE_SIZE),\n            T.RandomHorizontalFlip(),\n            T.ToTensor(),\n            T.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n        ])\n        self.transform = aug_tfm if augment else base_tfm\n\n        description = ({'image': 'byte', 'class': 'int'}\n                       if labeled else {'image': 'byte', 'id': 'byte'})\n\n        for path in filenames:\n            ds = TFRecordDataset(path, index_path=None, description=description)\n            for record in ds:\n                img_bytes = bytes(record['image'])\n                if labeled:\n                    self.samples.append((img_bytes, int(record['class'][0])))\n                else:\n                    self.samples.append((img_bytes, bytes(record['id']).decode('utf-8')))\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        img_bytes, target = self.samples[idx]\n        image = Image.open(io.BytesIO(img_bytes)).convert('RGB')\n        return self.transform(image), target","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Quick inspection of the training dataset","metadata":{}},{"cell_type":"code","source":"_inspect_ds = FlowerDataset(TRAINING_FILENAMES, labeled=True, augment=False)\nall_labels  = [lbl for _, lbl in _inspect_ds]\n\ncounter = Counter(all_labels)\n\nplt.figure(figsize=(20, 4))\nplt.bar(range(len(CLASSES)), [counter.get(i, 0) for i in range(len(CLASSES))])\nplt.xticks(range(len(CLASSES)), CLASSES, rotation=90, fontsize=7)\nplt.title('Class Distribution')\nplt.ylabel('Sample count')\nplt.tight_layout()\nplt.show()\n\ncounts = [counter.get(i, 0) for i in range(len(CLASSES))]\nprint(f'\\nMax/Min ratio: {max(counts)/min(counts):.1f}x')\nprint(f'Most common:  {CLASSES[np.argmax(counts)]} ({max(counts)})')\nprint(f'Least common: {CLASSES[np.argmin(counts)]} ({min(counts)})')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data pipelines","metadata":{}},{"cell_type":"code","source":"BATCH_SIZE = 64   # Lower to 16 if GPU memory is limited\n\ntrain_dataset = FlowerDataset(TRAINING_FILENAMES,   labeled=True,  augment=True)\nvalid_dataset = FlowerDataset(VALIDATION_FILENAMES, labeled=True,  augment=False)\ntest_dataset  = FlowerDataset(TEST_FILENAMES,        labeled=False, augment=False)\n\nds_train = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True,  num_workers=4, pin_memory=True)\nds_valid = DataLoader(valid_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=4, pin_memory=True)\nds_test  = DataLoader(test_dataset,  batch_size=BATCH_SIZE, shuffle=False, num_workers=4, pin_memory=True)\n\nNUM_TRAINING_IMAGES   = count_data_items(TRAINING_FILENAMES)\nNUM_VALIDATION_IMAGES = count_data_items(VALIDATION_FILENAMES)\nNUM_TEST_IMAGES       = count_data_items(TEST_FILENAMES)\nprint(f'Dataset: {NUM_TRAINING_IMAGES} training, {NUM_VALIDATION_IMAGES} validation, {NUM_TEST_IMAGES} test images')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Sanity-check shapes\nfor images, labels in ds_train:\n    print('Training batch — images:', images.shape, 'labels:', labels.shape)\n    break\nfor images, ids in ds_test:\n    print('Test batch    — images:', images.shape, 'sample ids:', ids[:3])\n    break","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train the model","metadata":{}},{"cell_type":"code","source":"EPOCHS    = 10\nFT_EPOCHS = 10","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model configuration","metadata":{}},{"cell_type":"code","source":"class ViTFlowerClassifier(nn.Module):\n    \"\"\"ViT backbone (HuggingFace AutoModel) with a lightweight classification head.\"\"\"\n\n    def __init__(self, num_classes, pretrained_name='google/vit-base-patch16-224'):\n        super().__init__()\n        self.backbone = AutoModel.from_pretrained(pretrained_name)\n        hidden_size   = self.backbone.config.hidden_size  # 768 for ViT-Base\n        self.head = nn.Sequential(\n            nn.Dropout(0.3),\n            nn.Linear(hidden_size, 256),\n            nn.GELU(),\n            nn.Dropout(0.2),\n            nn.Linear(256, num_classes),\n        )\n\n    def forward(self, pixel_values):\n        outputs   = self.backbone(pixel_values=pixel_values)\n        cls_token = outputs.last_hidden_state[:, 0, :]   # CLS token\n        return self.head(cls_token)\n\n    def freeze_backbone(self):\n        \"\"\"Freeze all backbone parameters.\"\"\"\n        for p in self.backbone.parameters():\n            p.requires_grad = False\n\n    def unfreeze_last_n_blocks(self, n):\n        self.freeze_backbone()\n        for layer in self.backbone.encoder.layer[-n:]:\n            for p in layer.parameters():\n                p.requires_grad = True\n        trainable = sum(p.numel() for p in self.parameters() if p.requires_grad)\n        print(f'Trainable parameters (last {n} blocks + head): {trainable:,}')\n\n\nmodel = ViTFlowerClassifier(num_classes=len(CLASSES)).to(DEVICE)\n\n# Phase 1: only the classification head is trainable\nmodel.freeze_backbone()\nfor p in model.head.parameters():\n    p.requires_grad = True\ntrainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f'Trainable parameters (head only): {trainable:,}')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training helpers","metadata":{}},{"cell_type":"code","source":"def run_epoch(model, loader, criterion, optimizer=None, scheduler=None, phase='train'):\n    \"\"\"Single epoch; returns (avg_loss, accuracy).\"\"\"\n    is_train = phase == 'train'\n    model.train(is_train)\n    total_loss, correct, total = 0.0, 0, 0\n\n    with torch.set_grad_enabled(is_train):\n        for images, labels in loader:\n            images = images.to(DEVICE, non_blocking=True)\n            labels = labels.to(DEVICE, non_blocking=True)\n\n            logits = model(images)\n            loss   = criterion(logits, labels)\n\n            if is_train:\n                optimizer.zero_grad()\n                loss.backward()\n                optimizer.step()\n                if scheduler is not None:\n                    scheduler.step()\n\n            total_loss += loss.item() * images.size(0)\n            correct    += (logits.argmax(1) == labels).sum().item()\n            total      += images.size(0)\n\n    return total_loss / total, correct / total\n\n\ndef train_model(model, ds_train, ds_valid, epochs, optimizer,\n               scheduler=None, patience=3, history=None):\n    \"\"\"Full training loop with early stopping. Appends to history if provided.\"\"\"\n    criterion = nn.CrossEntropyLoss()\n    if history is None:\n        history = {'train_loss': [], 'val_loss': [], 'train_acc': [], 'val_acc': []}\n\n    best_val_loss = float('inf')\n    patience_ctr  = 0\n    best_state    = None\n\n    for epoch in range(1, epochs + 1):\n        tr_loss, tr_acc = run_epoch(model, ds_train, criterion, optimizer, scheduler, 'train')\n        vl_loss, vl_acc = run_epoch(model, ds_valid, criterion, phase='val')\n\n        history['train_loss'].append(tr_loss)\n        history['val_loss'].append(vl_loss)\n        history['train_acc'].append(tr_acc)\n        history['val_acc'].append(vl_acc)\n\n        print(f'Epoch {epoch:02d}/{epochs}  '\n              f'train_loss={tr_loss:.4f}  train_acc={tr_acc:.4f}  '\n              f'val_loss={vl_loss:.4f}  val_acc={vl_acc:.4f}')\n\n        if vl_loss < best_val_loss:\n            best_val_loss = vl_loss\n            patience_ctr  = 0\n            best_state    = {k: v.cpu().clone() for k, v in model.state_dict().items()}\n        else:\n            patience_ctr += 1\n            if patience_ctr >= patience:\n                print('Early stopping triggered.')\n                break\n\n    if best_state is not None:\n        model.load_state_dict(best_state)\n        print('Restored best weights.')\n\n    return history\n\n\ndef plot_history(history, title='Training history'):\n    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(14, 4))\n    ax1.plot(history['train_loss'], label='train')\n    ax1.plot(history['val_loss'],   label='val')\n    ax1.set_title('Loss'); ax1.set_xlabel('Epoch'); ax1.legend()\n    ax2.plot(history['train_acc'], label='train')\n    ax2.plot(history['val_acc'],   label='val')\n    ax2.set_title('Accuracy'); ax2.set_xlabel('Epoch'); ax2.legend()\n    fig.suptitle(title)\n    plt.tight_layout()\n    plt.show()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Phase 1 — Train the head only","metadata":{}},{"cell_type":"code","source":"optimizer = optim.Adam(\n    filter(lambda p: p.requires_grad, model.parameters()), lr=1e-3\n)\n\nhistory = train_model(model, ds_train, ds_valid, epochs=EPOCHS,\n                      optimizer=optimizer, patience=3)\nplot_history(history, title='Head-only training')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Fine Tuning\n\nTo allow the model to learn more task-specific features we sequentially unfreeze the last N encoder blocks, each time continuing with a cosine-decay schedule and linear warm-up — mirroring the strategy from the original TensorFlow notebook.\n\nIf you want to unfreeze all layers, see the [HuggingFace documentation](https://huggingface.co/docs/transformers/training).","metadata":{}},{"cell_type":"code","source":"def cosine_schedule_with_warmup(optimizer, warmup_steps, total_steps, min_lr=1e-5):\n    \"\"\"LambdaLR scheduler: linear warm-up then cosine decay.\"\"\"\n    def lr_lambda(step):\n        if step < warmup_steps:\n            return step / max(1, warmup_steps)\n        progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)\n        cosine   = 0.5 * (1.0 + math.cos(math.pi * progress))\n        base_lr  = optimizer.param_groups[0]['initial_lr']\n        return max(min_lr / base_lr, cosine)\n    return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\n\n\nSTEPS_PER_EPOCH = len(ds_train)\n\n\ndef fine_tune(n_unfreeze, history_ft=None):\n    model.unfreeze_last_n_blocks(n_unfreeze)\n    total_steps  = FT_EPOCHS * STEPS_PER_EPOCH\n    warmup_steps = total_steps // 2\n    optimizer_ft = optim.Adam(\n        filter(lambda p: p.requires_grad, model.parameters()), lr=1e-5\n    )\n    for g in optimizer_ft.param_groups:\n        g['initial_lr'] = g['lr']\n    scheduler_ft = cosine_schedule_with_warmup(optimizer_ft, warmup_steps, total_steps)\n    return train_model(\n        model, ds_train, ds_valid,\n        epochs=FT_EPOCHS, optimizer=optimizer_ft,\n        scheduler=scheduler_ft, patience=3,\n        history=history_ft\n    )","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Unfreeze the last 2 blocks","metadata":{}},{"cell_type":"code","source":"history_ft = fine_tune(n_unfreeze=2)\nplot_history(history_ft, title='Fine-tuning — last 2 blocks')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Unfreeze the last 4 blocks","metadata":{}},{"cell_type":"code","source":"history_ft = fine_tune(n_unfreeze=4, history_ft=history_ft)\nplot_history(history_ft, title='Fine-tuning — last 4 blocks')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Unfreeze the last 6 blocks","metadata":{}},{"cell_type":"code","source":"history_ft = fine_tune(n_unfreeze=6, history_ft=history_ft)\nplot_history(history_ft, title='Fine-tuning — last 6 blocks')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Validation","metadata":{}},{"cell_type":"code","source":"plot_history(history,    title='Phase 1 — head only')\nplot_history(history_ft, title='Fine-tuning (all phases combined)')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix\n\nmodel.eval()\nall_preds, all_labels_val = [], []\n\nwith torch.no_grad():\n    for images, labels in ds_valid:\n        images = images.to(DEVICE)\n        logits = model(images)\n        all_preds.extend(logits.argmax(1).cpu().numpy())\n        all_labels_val.extend(labels.numpy())\n\nall_preds      = np.array(all_preds)\nall_labels_val = np.array(all_labels_val)\n\ncm = confusion_matrix(all_labels_val, all_preds)\nfig, ax = plt.subplots(figsize=(18, 16))\nim = ax.imshow(cm, interpolation='nearest', cmap='Blues')\nplt.colorbar(im, ax=ax)\nax.set(xticks=range(len(CLASSES)), yticks=range(len(CLASSES)),\n       xticklabels=CLASSES, yticklabels=CLASSES,\n       xlabel='Predicted', ylabel='True', title='Confusion Matrix')\nplt.setp(ax.get_xticklabels(), rotation=90, ha='right', fontsize=5)\nplt.setp(ax.get_yticklabels(), fontsize=5)\nplt.tight_layout()\nplt.show()\n\nprint(f'Validation accuracy: {(all_preds == all_labels_val).mean():.4f}')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Generate submission","metadata":{}},{"cell_type":"code","source":"model.eval()\nprob_list  = []\nid_list    = []\n\nprint('Computing predictions...')\nwith torch.no_grad():\n    for images, ids in ds_test:\n        images = images.to(DEVICE)\n        probs  = torch.softmax(model(images), dim=1)\n        prob_list.append(probs.cpu().numpy())\n        id_list.extend(ids)\n\nprobabilities = np.concatenate(prob_list, axis=0)\npredictions   = np.argmax(probabilities, axis=1)\nprint('Predictions shape:', predictions.shape)\n\nprint('Generating submission.csv file...')\npd.DataFrame({'id': id_list, 'label': predictions}).to_csv('submission.csv', index=False)\n\n# Preview\nimport subprocess\nsubprocess.run(['head', 'submission.csv'])","metadata":{},"outputs":[],"execution_count":null}]}