{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.0"},"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":"# LR Schedule Experiment : ViT Fine-Tuning for Flower Classification\n\nThis notebook investigates the influence of the learning rate schedule on fine-tuning a Vision Transformer (ViT) for the [Petals to the Metal](https://www.kaggle.com/c/tpu-getting-started) flower classification competition.\n\n**Experimental design:**\n- A single Phase 1 (head-only) training run produces a shared checkpoint.\n- Each experiment reloads that checkpoint and fine-tunes with the **last 4 encoder blocks unfrozen**.\n- The only variable between experiments is the LR schedule.\n\n**Schedules tested:**\n1. Baseline — constant LR (Adam, lr=1e-5)\n2. Cosine decay with warm-up\n3. Cosine decay without warm-up\n4. ReduceLROnPlateau (triggered on validation loss)\n\nA final cell overlays all runs for direct comparison.","metadata":{}},{"cell_type":"markdown","source":"# 1 Setup","metadata":{}},{"cell_type":"code","source":"!pip install transformers torch torchvision tfrecord scikit-learn","metadata":{},"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 copy\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":"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":"# 2 Dataset","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')\n\n# Extra training data\nGCS_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\nTRAINING_FILENAMES = (\n    TRAINING_FILENAMES\n    + list_tfrec(GCS_PATH_imagenet)\n    + list_tfrec(GCS_PATH_inaturalist)\n    + list_tfrec(GCS_PATH_openimage)\n    + list_tfrec(GCS_PATH_oxford_102)\n)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CLASSES = [\n    'pink primrose',    'hard-leaved pocket orchid', 'canterbury bells', 'sweet pea',     'wild geranium',     'tiger lily',           'moon orchid',              'bird of paradise', 'monkshood',        'globe thistle',\n    'snapdragon',       \"colt's foot\",               'king protea',      'spear thistle', 'yellow iris',       'globe-flower',         'purple coneflower',        'peruvian lily',    'balloon flower',   'giant white arum lily',\n    'fire lily',        'pincushion flower',         'fritillary',       'red ginger',    'grape hyacinth',    'corn poppy',           'prince of wales feathers', 'stemless gentian', 'artichoke',        'sweet william',\n    'carnation',        'garden phlox',              'love in the mist', 'cosmos',        'alpine sea holly',  'ruby-lipped cattleya', 'cape flower',              'great masterwort', 'siam tulip',       'lenten rose',\n    'barberton daisy',  'daffodil',                  'sword lily',       'poinsettia',    'bolero deep blue',  'wallflower',           'marigold',                 'buttercup',        'daisy',            'common dandelion',\n    'petunia',          'wild pansy',                'primula',          'sunflower',     'lilac hibiscus',    'bishop of llandaff',   'gaura',                    'geranium',         'orange dahlia',    'pink-yellow dahlia',\n    'cautleya spicata', 'japanese anemone',          'black-eyed susan', 'silverbush',    'californian poppy', 'osteospermum',         'spring crocus',            'iris',             'windflower',       'tree poppy',\n    'gazania',          'azalea',                    'water lily',       'rose',          'thorn apple',       'morning glory',        'passion flower',           'lotus',            'toad lily',        'anthurium',\n    'frangipani',       'clematis',                  'hibiscus',         'columbine',     'desert-rose',       'tree mallow',          'magnolia',                 'cyclamen ',        'watercress',       'canna lily',\n    'hippeastrum ',     'bee balm',                  'pink quill',       'foxglove',      'bougainvillea',     'camellia',             'mallow',                   'mexican petunia',  'bromelia',         'blanket flower',\n    'trumpet creeper',  'blackberry lily',           'common tulip',     'wild rose'\n]","metadata":{},"outputs":[],"execution_count":null},{"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    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    def __init__(self, filenames, labeled=True, augment=False):\n        self.labeled  = labeled\n        self.samples  = []\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":"code","source":"BATCH_SIZE = 64\n\ntrain_dataset = FlowerDataset(TRAINING_FILENAMES,   labeled=True, augment=True)\nvalid_dataset = FlowerDataset(VALIDATION_FILENAMES, labeled=True, 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)\n\nNUM_TRAINING_IMAGES   = count_data_items(TRAINING_FILENAMES)\nNUM_VALIDATION_IMAGES = count_data_items(VALIDATION_FILENAMES)\nprint(f'Dataset: {NUM_TRAINING_IMAGES} training, {NUM_VALIDATION_IMAGES} validation images')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3 Model & Training Helpers","metadata":{}},{"cell_type":"code","source":"class ViTFlowerClassifier(nn.Module):\n    \"\"\"ViT backbone 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, :]\n        return self.head(cls_token)\n\n    def freeze_backbone(self):\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:,}')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_epoch(model, loader, criterion, optimizer=None, scheduler=None,\n              phase='train', plateau_scheduler=None):\n    \"\"\"Single epoch. plateau_scheduler is stepped after validation, not per batch.\"\"\"\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:          # per-step schedulers\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    avg_loss = total_loss / total\n    accuracy = correct / total\n    return avg_loss, accuracy\n\n\ndef train_model(model, ds_train, ds_valid, epochs, optimizer,\n                scheduler=None, plateau_scheduler=None,\n                patience=3, history=None):\n    \"\"\"\n    Full training loop with early stopping.\n    - scheduler       : per-step LambdaLR (cosine variants)\n    - plateau_scheduler: ReduceLROnPlateau, stepped on val_loss after each epoch\n    \"\"\"\n    criterion = nn.CrossEntropyLoss()\n    if history is None:\n        history = {'train_loss': [], 'val_loss': [], 'train_acc': [], 'val_acc': [], 'lr': []}\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        # ReduceLROnPlateau steps on val loss (per epoch)\n        if plateau_scheduler is not None:\n            plateau_scheduler.step(vl_loss)\n\n        current_lr = optimizer.param_groups[0]['lr']\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        history['lr'].append(current_lr)\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              f'lr={current_lr:.2e}')\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, axes = plt.subplots(1, 3, figsize=(18, 4))\n    axes[0].plot(history['train_loss'], label='train')\n    axes[0].plot(history['val_loss'],   label='val')\n    axes[0].set_title('Loss'); axes[0].set_xlabel('Epoch'); axes[0].legend()\n    axes[1].plot(history['train_acc'], label='train')\n    axes[1].plot(history['val_acc'],   label='val')\n    axes[1].set_title('Accuracy'); axes[1].set_xlabel('Epoch'); axes[1].legend()\n    axes[2].plot(history['lr'])\n    axes[2].set_title('Learning Rate'); axes[2].set_xlabel('Epoch')\n    axes[2].set_yscale('log')\n    fig.suptitle(title)\n    plt.tight_layout()\n    plt.show()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def cosine_schedule_with_warmup(optimizer, warmup_steps, total_steps, min_ratio=1e-2):\n    \"\"\"Linear warm-up then cosine decay (per-step LambdaLR).\n    Returns a multiplier in [0, 1]; LambdaLR scales it against the optimizer's base LR.\n    min_ratio: floor as a fraction of the peak LR (default 1e-2 = 1% of peak).\n    \"\"\"\n    def lr_lambda(step):\n        if step < warmup_steps:\n            return step / max(1, warmup_steps)              # linear ramp 0 → 1\n        progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)\n        cosine   = 0.5 * (1.0 + math.cos(math.pi * progress))  # 1 → 0\n        return max(min_ratio, cosine)\n    return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\n\n\ndef cosine_schedule_no_warmup(optimizer, total_steps, min_ratio=1e-2):\n    \"\"\"Pure cosine decay from step 0, no warm-up (per-step LambdaLR).\n    Returns a multiplier in [0, 1]; LambdaLR scales it against the optimizer's base LR.\n    \"\"\"\n    def lr_lambda(step):\n        progress = step / max(1, total_steps)\n        cosine   = 0.5 * (1.0 + math.cos(math.pi * progress))  # 1 → 0\n        return max(min_ratio, cosine)\n    return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\n\n\ndef make_optimizer(model, lr=1e-5):\n    \"\"\"Adam over trainable parameters.\"\"\"\n    return optim.Adam(filter(lambda p: p.requires_grad, model.parameters()), lr=lr)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4 Phase 1: Head-Only Training (shared checkpoint)\n\nThis runs once. The resulting weights are saved to `phase1_checkpoint.pt` and reloaded at the start of every experiment so each fine-tuning run starts from an identical state.","metadata":{}},{"cell_type":"code","source":"EPOCHS    = 10\nFT_EPOCHS = 10\nFT_LR     = 1e-5\nN_UNFREEZE = 4   # fixed for all experiments\n\nCHECKPOINT_PATH = 'phase1_checkpoint.pt'","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = ViTFlowerClassifier(num_classes=len(CLASSES)).to(DEVICE)\n\n# Freeze backbone; only train the head\nmodel.freeze_backbone()\nfor p in model.head.parameters():\n    p.requires_grad = True\n\ntrainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f'Trainable parameters (head only): {trainable:,}')\n\noptimizer_p1 = optim.Adam(\n    filter(lambda p: p.requires_grad, model.parameters()), lr=1e-3\n)\n\nhistory_p1 = train_model(model, ds_train, ds_valid,\n                          epochs=EPOCHS, optimizer=optimizer_p1, patience=3)\n\ntorch.save(model.state_dict(), CHECKPOINT_PATH)\nprint(f'Phase 1 checkpoint saved to {CHECKPOINT_PATH}')\nplot_history(history_p1, title='Phase 1 — head only')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 5 Experiments\n\nEach cell below:\n1. Reloads the Phase 1 checkpoint into a fresh model instance.\n2. Unfreezes the last 4 encoder blocks.\n3. Fine-tunes with a specific LR schedule.\n4. Stores its history in `all_histories` for the final comparison.","metadata":{}},{"cell_type":"code","source":"# Container for all experiment results — populated by the cells below\nall_histories = {}","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Experiment 1: Baseline - Constant LR","metadata":{}},{"cell_type":"code","source":"model_exp1 = ViTFlowerClassifier(num_classes=len(CLASSES)).to(DEVICE)\nmodel_exp1.load_state_dict(torch.load(CHECKPOINT_PATH, map_location=DEVICE))\nmodel_exp1.unfreeze_last_n_blocks(N_UNFREEZE)\n\noptimizer_exp1 = make_optimizer(model_exp1, lr=FT_LR)\n\nhistory_exp1 = train_model(\n    model_exp1, ds_train, ds_valid,\n    epochs=FT_EPOCHS, optimizer=optimizer_exp1,\n    scheduler=None, plateau_scheduler=None,\n    patience=3\n)\n\nall_histories['Constant LR'] = history_exp1\nplot_history(history_exp1, title='Experiment 1 — Constant LR')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Experiment 2: Cosine Decay with Warm-Up","metadata":{}},{"cell_type":"code","source":"model_exp2 = ViTFlowerClassifier(num_classes=len(CLASSES)).to(DEVICE)\nmodel_exp2.load_state_dict(torch.load(CHECKPOINT_PATH, map_location=DEVICE))\nmodel_exp2.unfreeze_last_n_blocks(N_UNFREEZE)\n\noptimizer_exp2  = make_optimizer(model_exp2, lr=FT_LR)\ntotal_steps     = FT_EPOCHS * len(ds_train)\nwarmup_steps    = total_steps // 2\nscheduler_exp2  = cosine_schedule_with_warmup(optimizer_exp2, warmup_steps, total_steps)\n\nhistory_exp2 = train_model(\n    model_exp2, ds_train, ds_valid,\n    epochs=FT_EPOCHS, optimizer=optimizer_exp2,\n    scheduler=scheduler_exp2,\n    patience=3\n)\n\nall_histories['Cosine + Warm-Up'] = history_exp2\nplot_history(history_exp2, title='Experiment 2 — Cosine Decay with Warm-Up')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Experiment 3: Cosine Decay without Warm-Up","metadata":{}},{"cell_type":"code","source":"model_exp3 = ViTFlowerClassifier(num_classes=len(CLASSES)).to(DEVICE)\nmodel_exp3.load_state_dict(torch.load(CHECKPOINT_PATH, map_location=DEVICE))\nmodel_exp3.unfreeze_last_n_blocks(N_UNFREEZE)\n\noptimizer_exp3 = make_optimizer(model_exp3, lr=FT_LR)\ntotal_steps    = FT_EPOCHS * len(ds_train)\nscheduler_exp3 = cosine_schedule_no_warmup(optimizer_exp3, total_steps)\n\nhistory_exp3 = train_model(\n    model_exp3, ds_train, ds_valid,\n    epochs=FT_EPOCHS, optimizer=optimizer_exp3,\n    scheduler=scheduler_exp3,\n    patience=3\n)\n\nall_histories['Cosine No Warm-Up'] = history_exp3\nplot_history(history_exp3, title='Experiment 3 — Cosine Decay without Warm-Up')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Experiment 4: ReduceLROnPlateau","metadata":{}},{"cell_type":"code","source":"model_exp4 = ViTFlowerClassifier(num_classes=len(CLASSES)).to(DEVICE)\nmodel_exp4.load_state_dict(torch.load(CHECKPOINT_PATH, map_location=DEVICE))\nmodel_exp4.unfreeze_last_n_blocks(N_UNFREEZE)\n\noptimizer_exp4  = make_optimizer(model_exp4, lr=FT_LR)\nplateau_sched   = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer_exp4,\n    mode='min',       # monitor val_loss\n    factor=0.5,       # halve the LR on plateau\n    patience=2,       # wait 2 epochs before reducing\n    min_lr=1e-7\n)\n\nhistory_exp4 = train_model(\n    model_exp4, ds_train, ds_valid,\n    epochs=FT_EPOCHS, optimizer=optimizer_exp4,\n    scheduler=None, plateau_scheduler=plateau_sched,\n    patience=3\n)\n\nall_histories['ReduceLROnPlateau'] = history_exp4\nplot_history(history_exp4, title='Experiment 4 — ReduceLROnPlateau')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 6 Comparison\n\nAll four experiments overlaid on a single figure for direct comparison.","metadata":{}},{"cell_type":"code","source":"fig, axes = plt.subplots(1, 3, figsize=(20, 5))\n\ncolors = ['tab:blue', 'tab:orange', 'tab:green', 'tab:red']\n\nfor (name, hist), color in zip(all_histories.items(), colors):\n    epochs_run = range(1, len(hist['val_acc']) + 1)\n    axes[0].plot(epochs_run, hist['val_loss'],  color=color, label=name)\n    axes[1].plot(epochs_run, hist['val_acc'],   color=color, label=name)\n    axes[2].plot(epochs_run, hist['lr'],        color=color, label=name)\n\naxes[0].set_title('Validation Loss');     axes[0].set_xlabel('Epoch'); axes[0].legend()\naxes[1].set_title('Validation Accuracy'); axes[1].set_xlabel('Epoch'); axes[1].legend()\naxes[2].set_title('Learning Rate');       axes[2].set_xlabel('Epoch'); axes[2].set_yscale('log'); axes[2].legend()\n\nfig.suptitle('LR Schedule Comparison — Fine-Tuning (last 4 blocks)', fontsize=14)\nplt.tight_layout()\nplt.show()\n\n# Summary table\nprint(f\"\\n{'Schedule':<25} {'Best Val Acc':>12} {'Best Val Loss':>14} {'Epochs Run':>11}\")\nprint('-' * 65)\nfor name, hist in all_histories.items():\n    best_acc  = max(hist['val_acc'])\n    best_loss = min(hist['val_loss'])\n    n_epochs  = len(hist['val_acc'])\n    print(f\"{name:<25} {best_acc:>12.4f} {best_loss:>14.4f} {n_epochs:>11}\")","metadata":{},"outputs":[],"execution_count":null}]}