{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# ImageNet ResNet-18 Training\nCleaned and formatted version of the supplied notebook code. Training logic and hyperparameters are retained.","metadata":{}},{"cell_type":"code","source":"# !kaggle kernels output creoiqube/imagenet1k-resnet -p /kaggle/working\n\n# !zip -r output.zip *","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:57:26.764678Z","iopub.execute_input":"2026-10-03T08:57:26.764924Z","iopub.status.idle":"2026-10-03T08:57:26.768831Z","shell.execute_reply.started":"2026-10-03T08:57:26.764901Z","shell.execute_reply":"2026-10-03T08:57:26.768057Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !cat /kaggle/working/.virtual_documents/__notebook_source__.ipynb\n# !cat /kaggle/working/notebook6fb207366d.log","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:57:26.77046Z","iopub.execute_input":"2026-10-03T08:57:26.770959Z","iopub.status.idle":"2026-10-03T08:57:26.782359Z","shell.execute_reply.started":"2026-10-03T08:57:26.770935Z","shell.execute_reply":"2026-10-03T08:57:26.781855Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !cp -r /kaggle/input/datasets/creoiqube/check-point/. /kaggle/working","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:57:26.783181Z","iopub.execute_input":"2026-10-03T08:57:26.783436Z","iopub.status.idle":"2026-10-03T08:57:26.796322Z","shell.execute_reply.started":"2026-10-03T08:57:26.783405Z","shell.execute_reply":"2026-10-03T08:57:26.795488Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import csv\nimport hashlib\nimport math\nimport os\nimport pickle\nimport random\nimport time\nfrom datetime import datetime, timedelta\nfrom pathlib import Path\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom PIL import Image, ImageFile\n# from tqdm.auto import tqdm\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom torchvision.transforms import InterpolationMode\nimport torchvision.transforms as T\nfrom torch.utils.tensorboard import SummaryWriter","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:57:26.79746Z","iopub.execute_input":"2026-10-03T08:57:26.798313Z","iopub.status.idle":"2026-10-03T08:57:51.45827Z","shell.execute_reply.started":"2026-10-03T08:57:26.79828Z","shell.execute_reply":"2026-10-03T08:57:51.457713Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Hyper Parameters\nSEED = 42\nNUM_CLASSES = 1000\nIMAGE_SIZE = 224\nIMAGE_CHANNELS = 3\n\nGLOBAL_BATCH_SIZE = 512\nEPOCHS = 10\nCOMPLETED_EPOCHS = 0\nBATCH_LOG_EVERY = 100\n\nNUM_WORKERS = 4\nPREFETCH_FACTOR = 1\n\nBASE_LR = 0.01\nMIN_LR = 1e-6\nWEIGHT_DECAY = 1e-4\nMOMENTUM = 0.9\nLABEL_SMOOTHING = 0.1\n\nWARMUP_EPOCHS = 1\nWARMUP_START_FACTOR = 0.1\nCOSINE_CYCLE_EPOCHS = 3\nCOSINE_CYCLE_MULT = 1.0\n\nCHECKPOINT_EVERY_STEPS = 500\nRESUME = False\nNUM_GPUS = min(2, torch.cuda.device_count())\nassert torch.cuda.is_available() and NUM_GPUS == 2, 'This notebook is configured for Kaggle 2×T4.'\nassert GLOBAL_BATCH_SIZE % NUM_GPUS == 0\nBATCH_PER_GPU = GLOBAL_BATCH_SIZE // NUM_GPUS\nDEVICE = torch.device('cuda:0')\nprint('GPUs:', [torch.cuda.get_device_name(i) for i in range(NUM_GPUS)])\nprint(f'Global batch={GLOBAL_BATCH_SIZE} | per GPU={BATCH_PER_GPU} | workers={NUM_WORKERS}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:57:51.459112Z","iopub.execute_input":"2026-10-03T08:57:51.459924Z","iopub.status.idle":"2026-10-03T08:57:51.771154Z","shell.execute_reply.started":"2026-10-03T08:57:51.459893Z","shell.execute_reply":"2026-10-03T08:57:51.770508Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Reproducibility and runtime\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\ntorch.backends.cudnn.benchmark = True\ntorch.backends.cuda.matmul.allow_tf32 = True\ntorch.backends.cudnn.allow_tf32 = True\ntorch.set_num_threads(NUM_WORKERS)\n\n# Paths\nDATA_ROOT = Path('/kaggle/input/competitions/imagenet-object-localization-challenge')\nROOT = DATA_ROOT / 'ILSVRC/Data/CLS-LOC'\nTRAIN_DIR, VAL_DIR = ROOT/'train', ROOT/'val'\nVAL_CSV = DATA_ROOT/'LOC_val_solution.csv'\nOUTPUT_DIR = Path('/kaggle/working/imagenet_results')\nOUTPUT_DIR.mkdir(parents=True, exist_ok=True)\nINDEX_DIR = Path('/kaggle/input/datasets/yogeshwaranselvam521/resnet18-imagenet1k-index-files')\nif not INDEX_DIR.exists():\n    raise FileNotFoundError(f\"Index dataset not mounted: {INDEX_DIR}\")\nTRAIN_INDEX = INDEX_DIR/'train_index.pkl'\nVAL_INDEX = INDEX_DIR/'val_index.pkl'\nTENSORBOARD_DIR = OUTPUT_DIR / \"tensorboard\"\nTENSORBOARD_DIR.mkdir(parents=True, exist_ok=True)\n\nwriter = SummaryWriter(log_dir=str(TENSORBOARD_DIR))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:57:51.771934Z","iopub.execute_input":"2026-10-03T08:57:51.772469Z","iopub.status.idle":"2026-10-03T08:57:51.790883Z","shell.execute_reply.started":"2026-10-03T08:57:51.772442Z","shell.execute_reply":"2026-10-03T08:57:51.79041Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CPU PREPROCESSING\n\nIMAGENET_MEAN = (0.485, 0.456, 0.406)\nIMAGENET_STD = (0.229, 0.224, 0.225)\n\ntrain_transform = transforms.Compose([\n    transforms.RandomResizedCrop(IMAGE_SIZE, scale=(0.20, 1.0), ratio=(0.75, 1.333), interpolation=InterpolationMode.BILINEAR),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.PILToTensor(),\n])\n\neval_transform = transforms.Compose([\n    transforms.Resize(256, interpolation=InterpolationMode.BILINEAR,),\n    transforms.CenterCrop(IMAGE_SIZE),\n    transforms.PILToTensor(),\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:57:51.794754Z","iopub.execute_input":"2026-10-03T08:57:51.795479Z","iopub.status.idle":"2026-10-03T08:57:51.802835Z","shell.execute_reply.started":"2026-10-03T08:57:51.795441Z","shell.execute_reply":"2026-10-03T08:57:51.801958Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# DATASET\n\nclass ImageNetSamples(Dataset):\n    def __init__(self, samples, transform):\n        self.samples = samples\n        self.transform = transform\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        path, label = self.samples[idx]\n        with Image.open(path) as im:\n            image = self.transform(im.convert(\"RGB\"))\n        return image, int(label)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:57:51.80359Z","iopub.execute_input":"2026-10-03T08:57:51.804765Z","iopub.status.idle":"2026-10-03T08:57:51.810687Z","shell.execute_reply.started":"2026-10-03T08:57:51.804658Z","shell.execute_reply":"2026-10-03T08:57:51.809992Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# LOAD INDEX FILES\n\ndef load_indices():\n    if not TRAIN_INDEX.exists():\n        raise FileNotFoundError(f\"Training index not found: {TRAIN_INDEX}\")\n    if not VAL_INDEX.exists():\n        raise FileNotFoundError(f\"Validation index not found: {VAL_INDEX}\")\n    \n    with open(TRAIN_INDEX, \"rb\") as f:\n        train = pickle.load(f)\n    with open(VAL_INDEX, \"rb\") as f:\n        val = pickle.load(f)\n\n    if isinstance(train, dict):\n        train = train[\"samples\"]\n    if isinstance(val, dict):\n        val = val[\"samples\"]\n    if not train or not val:\n        raise RuntimeError(\"Loaded index is empty.\")\n\n    class_dirs = sorted(p for p in TRAIN_DIR.iterdir() if p.is_dir())\n    if len(class_dirs) != NUM_CLASSES:\n        raise RuntimeError(f\"Expected {NUM_CLASSES} classes, found {len(class_dirs)}\")\n\n    class_to_idx = {p.name: i for i, p in enumerate(class_dirs)}\n    labels = [label for _, label in train]\n    if min(labels) < 0 or max(labels) >= NUM_CLASSES:\n        raise ValueError(\"Train labels outside [0, 999]\")\n\n    counts = np.bincount(labels, minlength=NUM_CLASSES)\n\n    print(\n        f\"Train: {len(train):,} | \"\n        f\"Val: {len(val):,} | \"\n        f\"Classes: {NUM_CLASSES} | \"\n        f\"Missing: {(counts == 0).sum()}\"\n    )\n\n    return train, val, class_to_idx\n\n\ntrain_samples, val_samples, class_to_idx = load_indices()\n\ntrain_ds = ImageNetSamples(train_samples, train_transform)\nval_ds = ImageNetSamples(val_samples, eval_transform)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:57:51.812484Z","iopub.execute_input":"2026-10-03T08:57:51.813408Z","iopub.status.idle":"2026-10-03T08:57:56.535165Z","shell.execute_reply.started":"2026-10-03T08:57:51.813373Z","shell.execute_reply":"2026-10-03T08:57:56.534401Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# DATA LOADERS\n\ndef make_train_loader(epoch):\n    g = torch.Generator()\n    g.manual_seed(SEED + epoch)\n\n    return DataLoader(\n        train_ds, batch_size=GLOBAL_BATCH_SIZE, shuffle=True, generator=g,\n        num_workers=NUM_WORKERS, pin_memory=True, persistent_workers=(NUM_WORKERS > 0),\n        prefetch_factor=PREFETCH_FACTOR, drop_last=True,\n    )\n\n\nval_loader = DataLoader(\n    val_ds, batch_size=GLOBAL_BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS,\n    pin_memory=True, persistent_workers=(NUM_WORKERS > 0),\n    prefetch_factor=PREFETCH_FACTOR, drop_last=False,\n)\n\nprint(\"Steps/epoch:\", len(train_ds) // GLOBAL_BATCH_SIZE)\nprint(\"Validation batches:\", len(val_loader))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:57:56.536065Z","iopub.execute_input":"2026-10-03T08:57:56.536273Z","iopub.status.idle":"2026-10-03T08:57:56.543504Z","shell.execute_reply.started":"2026-10-03T08:57:56.536251Z","shell.execute_reply":"2026-10-03T08:57:56.542628Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# GPU PREPROCESSING\n\nGPU_MEAN = torch.tensor(IMAGENET_MEAN, device=DEVICE, dtype=torch.float32).view(1, 3, 1, 1)\nGPU_STD = torch.tensor(IMAGENET_STD, device=DEVICE, dtype=torch.float32).view(1, 3, 1, 1)\n\ngpu_train_augmentation = T.Compose([\n    T.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.05,),\n    T.RandomErasing(p=0.25, scale=(0.02, 0.20), ratio=(0.3, 3.3), value=0,),\n])\n\ndef preprocess_gpu(images, training=True):\n    images = images.to(DEVICE, non_blocking=True, memory_format=torch.channels_last)\n    images = images.float().div_(255.0)\n\n    if training:\n        images = gpu_train_augmentation(images)\n    images.sub_(GPU_MEAN).div_(GPU_STD)\n    images = images.to(dtype=torch.float16, memory_format=torch.channels_last)\n\n    return images","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:57:56.544494Z","iopub.execute_input":"2026-10-03T08:57:56.544835Z","iopub.status.idle":"2026-10-03T08:57:56.837141Z","shell.execute_reply.started":"2026-10-03T08:57:56.544797Z","shell.execute_reply":"2026-10-03T08:57:56.836256Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Residual Block\n\nclass ResidualBlock(nn.Module):\n\n    def __init__(self, in_channels, out_channels, stride=1):\n        super().__init__()\n\n        self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1, bias=False)\n        self.bn1 = nn.BatchNorm2d(out_channels)\n        self.relu = nn.ReLU(inplace=True)\n        self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False)\n        self.bn2 = nn.BatchNorm2d(out_channels)\n\n        if stride != 1 or in_channels != out_channels:\n            self.shortcut = nn.Sequential(\n                nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride, bias=False),\n                nn.BatchNorm2d(out_channels)\n            )\n\n        else:\n            self.shortcut = nn.Identity()\n\n    def forward(self, x):\n\n        identity = self.shortcut(x)\n\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n\n        out = self.conv2(out)\n        out = self.bn2(out)\n\n        out += identity\n\n        out = self.relu(out)\n\n        return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:57:56.838222Z","iopub.execute_input":"2026-10-03T08:57:56.838573Z","iopub.status.idle":"2026-10-03T08:57:56.845919Z","shell.execute_reply.started":"2026-10-03T08:57:56.838533Z","shell.execute_reply":"2026-10-03T08:57:56.845314Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ResNet-18\n\nclass ResNet18(nn.Module):\n\n    def __init__(self, num_classes=NUM_CLASSES):\n        super().__init__()\n\n        self.in_channels = 64\n\n        self.conv1 = nn.Conv2d(in_channels=IMAGE_CHANNELS, out_channels=64, kernel_size=7, stride=2, padding=3, bias=False)\n        self.bn1 = nn.BatchNorm2d(64)\n        self.relu = nn.ReLU(inplace=True)\n        self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)\n\n        self.layer1 = self._make_layer(out_channels=64, num_blocks=2, stride=1)\n        self.layer2 = self._make_layer(out_channels=128, num_blocks=2, stride=2)\n        self.layer3 = self._make_layer(out_channels=256, num_blocks=2, stride=2)\n        self.layer4 = self._make_layer(out_channels=512, num_blocks=2, stride=2)\n        self.avgpool = nn.AdaptiveAvgPool2d((1, 1))\n\n        self.fc = nn.Linear(512, num_classes)\n\n    def _make_layer(self, out_channels, num_blocks, stride):\n        strides = [stride] + [1] * (num_blocks - 1)\n        layers = []\n\n        for stride in strides:\n            layers.append(ResidualBlock(self.in_channels, out_channels, stride))\n\n            self.in_channels = out_channels\n\n        return nn.Sequential(*layers)\n\n    def forward(self, x):\n        x = self.conv1(x)\n        x = self.bn1(x)\n        x = self.relu(x)\n        x = self.maxpool(x)\n\n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n        x = self.avgpool(x)\n\n        x = torch.flatten(x, 1)\n        x = self.fc(x)\n\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:57:56.846941Z","iopub.execute_input":"2026-10-03T08:57:56.847291Z","iopub.status.idle":"2026-10-03T08:57:56.866171Z","shell.execute_reply.started":"2026-10-03T08:57:56.847269Z","shell.execute_reply":"2026-10-03T08:57:56.865438Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_lr(step, base_lr):\n    warmup_steps = WARMUP_EPOCHS * steps_per_epoch\n\n    # 1. Linear warmup\n    if step < warmup_steps:\n        warmup_progress = step / max(1, warmup_steps)\n        return base_lr * (WARMUP_START_FACTOR + (1.0 - WARMUP_START_FACTOR) * warmup_progress)\n\n    # 2. Cosine annealing with restarts\n    cosine_step = step - warmup_steps\n    cycle_steps = COSINE_CYCLE_EPOCHS * steps_per_epoch\n    cycle_position = (cosine_step % cycle_steps) / cycle_steps\n    cosine_factor = 0.5 * (1.0 + math.cos(math.pi * cycle_position))\n    \n    return MIN_LR + (base_lr - MIN_LR) * cosine_factor\n\ndef update_lr(optimizer, step, base_lr):\n    lr = get_lr(step, base_lr)\n    for param_group in optimizer.param_groups:\n        param_group[\"lr\"] = lr\n    return lr","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:57:56.867201Z","iopub.execute_input":"2026-10-03T08:57:56.867599Z","iopub.status.idle":"2026-10-03T08:57:56.882326Z","shell.execute_reply.started":"2026-10-03T08:57:56.867571Z","shell.execute_reply":"2026-10-03T08:57:56.881544Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Model, loss, optimizer, and learning-rate schedule\nbase_model = ResNet18(num_classes=NUM_CLASSES).to(DEVICE,memory_format=torch.channels_last)\nmodel = nn.DataParallel(base_model, device_ids=[0, 1], output_device=0)\ncriterion = nn.CrossEntropyLoss(label_smoothing=LABEL_SMOOTHING)\n\ndecay = []\nno_decay = []\n\nfor name, p in base_model.named_parameters():\n    if not p.requires_grad: continue\n    (decay if p.ndim > 1 else no_decay).append(p)\n\noptimizer = torch.optim.SGD(\n    [\n        {\"params\": decay, \"weight_decay\": WEIGHT_DECAY},\n        {\"params\": no_decay, \"weight_decay\": 0.0},\n    ],\n    lr=BASE_LR, momentum=0.0, nesterov=False,\n)\nscaler = torch.amp.GradScaler(\"cuda\", enabled=True)\n\nprint(\"Optimizer: Vanilla SGD\")\nprint(\"Base LR:\", BASE_LR)\nprint(\"Momentum:\", 0.0)\n\nsteps_per_epoch = len(train_ds) // GLOBAL_BATCH_SIZE\nremaining_steps = steps_per_epoch * (EPOCHS - COMPLETED_EPOCHS)\n\nprint('Trainable parameters:',sum(p.numel() for p in base_model.parameters() if p.requires_grad))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:57:56.88348Z","iopub.execute_input":"2026-10-03T08:57:56.884151Z","iopub.status.idle":"2026-10-03T08:57:57.067013Z","shell.execute_reply.started":"2026-10-03T08:57:56.884114Z","shell.execute_reply":"2026-10-03T08:57:57.066301Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# PIPELINE VERIFICATION\n\ndef verify_pipeline(name, loader, training):\n    images, labels = next(iter(loader))\n\n    print(f\"\\n{name} | CPU: {images.shape}, {images.dtype}\")\n\n    assert images.shape[1:] == (3, IMAGE_SIZE, IMAGE_SIZE)\n    assert images.dtype == torch.uint8\n    assert labels.dtype == torch.int64\n\n    images = preprocess_gpu(images, training=training)\n    labels = labels.to(DEVICE, non_blocking=True)\n\n    print(f\"GPU: {images.shape}, {images.dtype}, {images.device}\")\n    print(f\"Channels Last: {images.is_contiguous(memory_format=torch.channels_last)}\")\n\n    assert images.dtype == torch.float16\n    assert images.device.type == \"cuda\"\n    assert images.is_contiguous(memory_format=torch.channels_last)\n    assert torch.isfinite(images).all()\n\n    model.eval()\n    with torch.inference_mode(), torch.autocast(\"cuda\", dtype=torch.float16):\n        logits = model(images)\n\n    print(f\"Logits: {logits.shape}, {logits.dtype}\")\n    assert logits.shape == (images.size(0), NUM_CLASSES)\n    assert torch.isfinite(logits).all()\n\n    print(f\"{name}: PASSED\")\n\n    del images, labels, logits\n    torch.cuda.empty_cache()\n\n\nverify_pipeline(\"TRAIN\", make_train_loader(0), True)\nverify_pipeline(\"VALIDATION\", val_loader, False)\n\nprint(\"\\nALL PIPELINE TESTS PASSED\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:57:57.067872Z","iopub.execute_input":"2026-10-03T08:57:57.068082Z","iopub.status.idle":"2026-10-03T08:58:22.959777Z","shell.execute_reply.started":"2026-10-03T08:57:57.06806Z","shell.execute_reply":"2026-10-03T08:58:22.958582Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"LAST_CKPT=OUTPUT_DIR/'last_checkpoint.pth'\nBEST_CKPT=OUTPUT_DIR/'best_checkpoint.pth'\nHISTORY_PATH=OUTPUT_DIR/'training_history.csv'\nhistory=[]\nbest_val_top1=0.0\nstart_epoch=0\nresume_batch=0\nglobal_step=COMPLETED_EPOCHS * steps_per_epoch\n\nif RESUME and HISTORY_PATH.exists():\n    try: history=pd.read_csv(HISTORY_PATH).to_dict('records')\n    except Exception: history=[]\nif RESUME and LAST_CKPT.exists():\n    ck=torch.load(LAST_CKPT,map_location=DEVICE,weights_only=False)\n    base_model.load_state_dict(ck['model'])\n    optimizer.load_state_dict(ck['optimizer'])\n    scaler.load_state_dict(ck['scaler'])\n    start_epoch=int(ck['epoch'])\n    resume_batch=int(ck.get('next_batch',0))\n    global_step=int(ck.get('global_step',0))\n    best_val_top1=float(ck.get('best_val_top1',0.0))\n    if resume_batch >= steps_per_epoch:\n        start_epoch+=1\n        resume_batch=0\n    print(f'Resumed at epoch index={start_epoch}, batch={resume_batch}, global_step={global_step}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:58:22.961926Z","iopub.execute_input":"2026-10-03T08:58:22.962311Z","iopub.status.idle":"2026-10-03T08:58:22.971433Z","shell.execute_reply.started":"2026-10-03T08:58:22.962267Z","shell.execute_reply":"2026-10-03T08:58:22.970543Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def save_state(epoch,next_batch,best):\n    state={'model':base_model.state_dict(),'optimizer':optimizer.state_dict(),'scaler':scaler.state_dict(),\n           'epoch':epoch,'next_batch':next_batch,'global_step':global_step,'best_val_top1':best,\n           'class_to_idx':class_to_idx,'config':{'epochs':EPOCHS,'global_batch':GLOBAL_BATCH_SIZE,'image_size':IMAGE_SIZE}}\n    tmp=Path(str(LAST_CKPT)+'.tmp')\n    torch.save(state,tmp)\n    os.replace(tmp,LAST_CKPT)\n\n\ndef topk_counts(logits, targets):\n    pred=logits.topk(5,dim=1).indices\n    top1=pred[:,0].eq(targets).sum()\n    top5=pred.eq(targets[:,None]).any(dim=1).sum()\n    return top1,top5","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:58:22.972698Z","iopub.execute_input":"2026-10-03T08:58:22.973006Z","iopub.status.idle":"2026-10-03T08:58:23.027581Z","shell.execute_reply.started":"2026-10-03T08:58:22.972972Z","shell.execute_reply":"2026-10-03T08:58:23.02673Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ck = torch.load(\n#     LAST_CKPT,\n#     map_location=DEVICE,\n#     weights_only=False\n# )\n\n# print(\"Checkpoint epoch:\", ck[\"epoch\"])\n# print(\"Checkpoint batch:\", ck.get(\"next_batch\"))\n# print(\"Global step:\", ck.get(\"global_step\"))\n# print(\"Best validation Top-1:\", ck.get(\"best_val_top1\"))\n\n# assert ck[\"epoch\"] == 2, (\n#     \"Expected a checkpoint after Epoch 2.\"\n# )\n\n# assert ck.get(\"next_batch\", 0) == 0, (\n#     \"Checkpoint is not at an epoch boundary.\"\n# )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:58:23.031798Z","iopub.execute_input":"2026-10-03T08:58:23.032024Z","iopub.status.idle":"2026-10-03T08:58:23.047261Z","shell.execute_reply.started":"2026-10-03T08:58:23.032003Z","shell.execute_reply":"2026-10-03T08:58:23.046468Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(epoch):\n    global global_step\n\n    model.train()\n    loader = make_train_loader(epoch)\n    skip = resume_batch if epoch == start_epoch else 0\n    loss_sum = 0.0\n    top1_sum = 0\n    top5_sum = 0\n    seen = 0\n    epoch_start = time.perf_counter()\n    batches_total = len(loader) - skip\n    # bar = tqdm(enumerate(loader), total=len(loader), desc=f\"Epoch {epoch + 1}/{EPOCHS}\")\n\n    for batch_idx, (images, targets) in enumerate(loader): #bar:\n        if batch_idx < skip:\n            continue\n        \n        images = preprocess_gpu(images, training=True)   \n        targets = targets.to(DEVICE, non_blocking=True)\n        current_lr = update_lr(optimizer, global_step, BASE_LR)\n        optimizer.zero_grad(set_to_none=True)\n\n        with torch.autocast(\"cuda\", dtype=torch.float16, enabled=True):\n            logits = model(images)\n            loss = criterion(logits, targets)\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        with torch.no_grad():\n            t1, t5 = topk_counts(logits.detach(), targets)\n        n = targets.size(0)\n        loss_sum += loss.detach().item() * n\n        top1_sum += t1.item()\n        top5_sum += t5.item()\n        seen += n\n        global_step += 1\n\n        if (batch_idx + 1) % BATCH_LOG_EVERY == 0:\n            elapsed = time.perf_counter() - epoch_start\n            batches_done = batch_idx + 1 - skip\n            avg_batch_time = elapsed / max(1, batches_done)\n            epoch_eta = avg_batch_time * (batches_total - batches_done)\n            remaining_epochs = EPOCHS - epoch - 1\n            total_eta = (epoch_eta + remaining_epochs * avg_batch_time * len(loader))\n            completion_time = datetime.now() + timedelta(seconds=total_eta)\n\n            print(\n                f\"\\nEpoch {epoch + 1}/{EPOCHS} | \"\n                f\"Batch {batch_idx + 1}/{len(loader)} | \"\n                f\"Loss: {loss.item():.4f} | \"\n                f\"Top-1: {100 * top1_sum / max(1, seen):.2f}% | \"\n                f\"Top-5: {100 * top5_sum / max(1, seen):.2f}% | \"\n                f\"LR: {current_lr:.6g} | \"\n                f\"Epoch ETA: {epoch_eta / 60:.1f} min | \"\n                f\"Total ETA: {total_eta / 3600:.2f} hrs | \"\n                f\"Finish: {completion_time.strftime('%d-%m %I:%M %p')}\",\n                flush=True\n            )\n\n            writer.add_scalar(\"Batch/Loss\", loss.item(), global_step)\n            writer.add_scalar(\"Batch/Top1_Accuracy\", 100 * t1.item() / n, global_step)\n            writer.add_scalar(\"Batch/Top5_Accuracy\", 100 * t5.item() / n, global_step)\n            writer.add_scalar(\"Batch/Learning_Rate\", current_lr, global_step)\n\n        # bar.set_postfix(loss=f\"{loss.item():.3f}\", top1=f\"{100 * top1_sum / max(1, seen):.2f}%\")\n\n        if global_step % CHECKPOINT_EVERY_STEPS == 0:\n            save_state(epoch, batch_idx + 1, best_val_top1)\n\n    return (loss_sum / max(1, seen), 100 * top1_sum / max(1, seen), 100 * top5_sum / max(1, seen), time.perf_counter() - epoch_start)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:58:23.048365Z","iopub.execute_input":"2026-10-03T08:58:23.048703Z","iopub.status.idle":"2026-10-03T08:58:23.070858Z","shell.execute_reply.started":"2026-10-03T08:58:23.04864Z","shell.execute_reply":"2026-10-03T08:58:23.070105Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef validate():\n\n    model.eval()\n\n    loss_sum = 0.0\n    top1_sum = 0\n    top5_sum = 0\n    total = 0\n\n    with torch.inference_mode():\n        for images, targets in val_loader: # tqdm(val_loader, desc=\"Validation\", leave=False):\n            images = preprocess_gpu(images, training=False)   \n            targets = targets.to(DEVICE, non_blocking=True)\n    \n            with torch.autocast(\"cuda\", dtype=torch.float16, enabled=True):\n                logits = model(images)\n            loss = nn.functional.cross_entropy(logits, targets, reduction=\"sum\")\n            t1, t5 = topk_counts(logits, targets)\n            n = targets.size(0)\n            loss_sum += loss.item()\n            top1_sum += t1.item()\n            top5_sum += t5.item()\n            total += n\n\n    return (loss_sum / total, 100 * top1_sum / total, 100 * top5_sum / total)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:58:23.071869Z","iopub.execute_input":"2026-10-03T08:58:23.073125Z","iopub.status.idle":"2026-10-03T08:58:23.090241Z","shell.execute_reply.started":"2026-10-03T08:58:23.07309Z","shell.execute_reply":"2026-10-03T08:58:23.089599Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate_epoch(epoch, train_metrics, epoch_seconds):\n    global best_val_top1, history\n\n    train_loss, train_top1, train_top5 = train_metrics\n    val_loss, val_top1, val_top5 = validate()\n\n    writer.add_scalars(\"Epoch/Loss\", {\"train\": train_loss, \"validation\": val_loss}, epoch + 1)\n    writer.add_scalars(\"Epoch/Top1_Accuracy\", {\"train\": train_top1, \"validation\": val_top1}, epoch + 1)\n    writer.add_scalars(\"Epoch/Top5_Accuracy\", {\"train\": train_top5, \"validation\": val_top5}, epoch + 1)\n\n    previous_best = best_val_top1\n    best_val_top1 = max(best_val_top1, val_top1)\n\n    row = {\n        \"epoch\": epoch + 1,\n        \"train_loss\": train_loss,\n        \"train_top1\": train_top1,\n        \"train_top5\": train_top5,\n        \"val_loss\": val_loss,\n        \"val_top1\": val_top1,\n        \"val_top5\": val_top5,\n        \"lr_end\": lr_for_step(max(0, global_step - 1)),\n        \"epoch_seconds\": epoch_seconds\n    }\n\n    history = [r for r in history if int(r.get(\"epoch\", -1)) != epoch + 1]\n    history.append(row)\n    pd.DataFrame(history).sort_values(\"epoch\").to_csv(HISTORY_PATH, index=False)\n    save_state(epoch + 1, 0, best_val_top1)\n\n    if val_top1 >= previous_best:\n        torch.save({\n            \"model\": base_model.state_dict(),\n            \"epoch\": epoch + 1,\n            \"val_top1\": val_top1,\n            \"class_to_idx\": class_to_idx\n        }, BEST_CKPT)\n\n    print(\n        f\"\\nEpoch {epoch + 1}/{EPOCHS} COMPLETE\\n\"\n        f\"Train Loss: {train_loss:.4f} | \"\n        f\"Train Top-1: {train_top1:.2f}% | \"\n        f\"Train Top-5: {train_top5:.2f}%\\n\"\n        f\"Val Loss: {val_loss:.4f} | \"\n        f\"Val Top-1: {val_top1:.2f}% | \"\n        f\"Val Top-5: {val_top5:.2f}%\\n\"\n        f\"Training Duration: {epoch_seconds / 60:.1f} minutes\\n\",\n        flush=True\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:58:23.092036Z","iopub.execute_input":"2026-10-03T08:58:23.092396Z","iopub.status.idle":"2026-10-03T08:58:23.107391Z","shell.execute_reply.started":"2026-10-03T08:58:23.092358Z","shell.execute_reply":"2026-10-03T08:58:23.10651Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Training Phase\n\nfor epoch in range(start_epoch, EPOCHS):\n    train_metrics = train_one_epoch(epoch)\n    evaluate_epoch(epoch, train_metrics[:3], train_metrics[3])\n    resume_batch = 0\n\nwriter.flush()\nwriter.close()\nprint(f\"Training finished. Last: {LAST_CKPT} | Best: {BEST_CKPT} | History: {HISTORY_PATH}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-10-03T08:58:23.108408Z","iopub.execute_input":"2026-10-03T08:58:23.108802Z","execution_failed":"2026-10-03T08:59:24.066Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plot epoch-level classification metrics\nif HISTORY_PATH.exists():\n    h=pd.read_csv(HISTORY_PATH).sort_values('epoch')\n    fig,ax=plt.subplots(1,2,figsize=(12,4))\n    ax[0].plot(h.epoch,h.train_loss,label='train')\n    ax[0].plot(h.epoch,h.val_loss,label='validation')\n    ax[0].set_ylabel('Cross-entropy')\n    ax[0].legend()\n    ax[0].grid(alpha=.3)\n    ax[1].plot(h.epoch,h.train_top1,label='train top-1')\n    ax[1].plot(h.epoch,h.val_top1,label='val top-1')\n    ax[1].plot(h.epoch,h.train_top5,label='train top-5')\n    ax[1].plot(h.epoch,h.val_top5,label='val top-5')\n    ax[1].set_ylabel('Accuracy (%)')\n    ax[1].legend()\n    ax[1].grid(alpha=.3)\n    for a in ax: a.set_xlabel('Epoch')\n    plt.tight_layout()\n    plt.savefig(OUTPUT_DIR/'training_curves.png',dpi=160)\n    plt.show()","metadata":{"trusted":true,"execution":{"execution_failed":"2026-10-03T08:59:24.067Z"}},"outputs":[],"execution_count":null}]}