{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q webdataset\n\nimport os, io, random, tarfile, time, glob\nimport numpy as np, pandas as pd\nfrom PIL import Image\nfrom multiprocessing import Pool\n\nROOT  = \"/kaggle/input/competitions/imagenet-object-localization-challenge\"\nTRAIN = f\"{ROOT}/ILSVRC/Data/CLS-LOC/train\"\nVAL   = f\"{ROOT}/ILSVRC/Data/CLS-LOC/val\"\nOUT   = \"/kaggle/working/shards\"\nos.makedirs(OUT, exist_ok=True)\n\nprint(os.path.exists(TRAIN), os.path.exists(VAL), os.path.exists(f\"{ROOT}/LOC_val_solution.csv\"))\n\nclasses = sorted(os.listdir(TRAIN))\ncls_to_idx = {c: i for i, c in enumerate(classes)}\nprint(len(classes))  # 1000","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-09-29T01:13:40.95263Z","iopub.execute_input":"2026-09-29T01:13:40.95345Z","iopub.status.idle":"2026-09-29T01:13:44.43987Z","shell.execute_reply.started":"2026-09-29T01:13:40.953412Z","shell.execute_reply":"2026-09-29T01:13:44.43905Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"QUALITY = 85  # lowers size a bit so everything fits under the 20 GB limit\n\ndef process(path):\n    img = Image.open(path)\n    img.draft(\"RGB\", (256, 256))              # fast partial JPEG decode\n    img = img.convert(\"RGB\")\n    w, h = img.size; s = 256 / min(w, h)\n    img = img.resize((round(w*s), round(h*s)), Image.BILINEAR)\n    w, h = img.size; l, t = (w-224)//2, (h-224)//2\n    img = img.crop((l, t, l+224, t+224))\n    buf = io.BytesIO(); img.save(buf, \"JPEG\", quality=QUALITY)\n    return buf.getvalue()\n\ndef write_shard(args):\n    shard_path, items, start = args\n    if os.path.exists(shard_path):            # resume: skip finished shards\n        return shard_path\n    tmp = shard_path + \".tmp\"\n    with tarfile.open(tmp, \"w\") as tar:\n        for k, (path, label) in enumerate(items):\n            try:\n                data = process(path)\n            except Exception as e:\n                print(\"skip\", path, e); continue\n            key = f\"{start + k:08d}\"\n            for ext, d in ((\"jpg\", data), (\"cls\", str(label).encode())):\n                info = tarfile.TarInfo(f\"{key}.{ext}\"); info.size = len(d)\n                tar.addfile(info, io.BytesIO(d))\n    os.rename(tmp, shard_path)                # only complete shards get the final name\n    return shard_path\n\ndef build(items, prefix, per_shard=10000):\n    jobs = [(f\"{OUT}/{prefix}-{i//per_shard:04d}.tar\", items[i:i+per_shard], i)\n            for i in range(0, len(items), per_shard)]\n    t0 = time.time()\n    with Pool(os.cpu_count()) as p:\n        for n, done in enumerate(p.imap_unordered(write_shard, jobs), 1):\n            print(f\"{n}/{len(jobs)} {os.path.basename(done)}  {(time.time()-t0)/60:.1f} min\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T08:38:59.81098Z","iopub.execute_input":"2026-09-29T08:38:59.811355Z","iopub.status.idle":"2026-09-29T08:38:59.823552Z","shell.execute_reply.started":"2026-09-29T08:38:59.811324Z","shell.execute_reply":"2026-09-29T08:38:59.822462Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"build(val_items, \"val\")\nbuild(train_items, \"train\")\n!du -sh /kaggle/working/shards","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T01:16:19.45985Z","iopub.execute_input":"2026-09-29T01:16:19.460393Z","iopub.status.idle":"2026-09-29T02:50:00.929943Z","shell.execute_reply.started":"2026-09-29T01:16:19.460307Z","shell.execute_reply":"2026-09-29T02:50:00.929039Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q webdataset\nimport os, io, time, glob\nimport numpy as np\nimport torch, torch.nn as nn\nimport webdataset as wds\n\nOUT = \"/kaggle/working/shards\"\ntrain_urls = sorted(glob.glob(f\"{OUT}/train-*.tar\"))\nval_urls   = sorted(glob.glob(f\"{OUT}/val-*.tar\"))\nprint(len(train_urls), len(val_urls), glob.glob(f\"{OUT}/*.tmp\"))   # expect 129 5 []\nN_TRAIN = 1281167\ndev = \"cuda\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T03:56:35.175438Z","iopub.execute_input":"2026-09-29T03:56:35.176291Z","iopub.status.idle":"2026-09-29T03:56:38.726652Z","shell.execute_reply.started":"2026-09-29T03:56:35.176252Z","shell.execute_reply":"2026-09-29T03:56:38.725654Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def to_tensor(img):\n    return torch.from_numpy(np.asarray(img).copy()).permute(2, 0, 1)\n\ndef make_loader(urls, train, bs=256, workers=4):\n    ds = wds.WebDataset(urls, shardshuffle=100 if train else False,\n                        handler=wds.warn_and_continue)\n    if train:\n        ds = ds.shuffle(2000)\n    ds = ds.decode(\"pil\").to_tuple(\"jpg\", \"cls\").map_tuple(to_tensor, int)\n    return torch.utils.data.DataLoader(ds, batch_size=bs, num_workers=workers,\n                                       pin_memory=True, drop_last=train,\n                                       prefetch_factor=4, persistent_workers=True)\n\nBS = 256\ntrain_loader = make_loader(train_urls, True, BS)\nval_loader   = make_loader(val_urls, False, 250)\n\nmean = torch.tensor([0.485,0.456,0.406], device=dev).view(1,3,1,1) * 255\nstd  = torch.tensor([0.229,0.224,0.225], device=dev).view(1,3,1,1) * 255\n\ndef prep(x):\n    return ((x.float() - mean) / std).contiguous(memory_format=torch.channels_last)\n\ndef gpu_flip(x):\n    m = torch.rand(x.size(0), device=x.device) < 0.5\n    return torch.where(m.view(-1,1,1,1), x.flip(3), x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T03:59:49.774321Z","iopub.execute_input":"2026-09-29T03:59:49.774725Z","iopub.status.idle":"2026-09-29T03:59:49.785131Z","shell.execute_reply.started":"2026-09-29T03:59:49.774699Z","shell.execute_reply":"2026-09-29T03:59:49.784096Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc, time\ntry:\n    del train_loader, val_loader\nexcept NameError:\n    pass\ngc.collect()\n\ndef make_loader(urls, train, bs=256, workers=3):\n    ds = wds.WebDataset(urls, shardshuffle=100 if train else False,\n                        handler=wds.warn_and_continue)\n    if train:\n        ds = ds.shuffle(1000)\n    ds = ds.decode(\"pil\").to_tuple(\"jpg\", \"cls\").map_tuple(to_tensor, int)\n    return torch.utils.data.DataLoader(ds, batch_size=bs, num_workers=workers,\n                                       pin_memory=True, drop_last=train,\n                                       prefetch_factor=2, persistent_workers=False)\n\nBS = 256\ntrain_loader = make_loader(train_urls, True, BS, workers=3)\nval_loader   = make_loader(val_urls, False, 250, workers=2)\n\nt = time.time()\nfor i, (x, y) in enumerate(train_loader):\n    if i == 20: break\nprint(x.shape, f\"{20*BS/(time.time()-t):.0f} img/s from loader alone\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T06:19:51.309593Z","iopub.execute_input":"2026-09-29T06:19:51.309826Z","iopub.status.idle":"2026-09-29T06:19:55.158852Z","shell.execute_reply.started":"2026-09-29T06:19:51.309802Z","shell.execute_reply":"2026-09-29T06:19:55.157793Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q webdataset\nimport os, io, time, glob, csv, gc, math\nimport numpy as np\nimport torch, torch.nn as nn\nimport webdataset as wds\nimport matplotlib.pyplot as plt\nfrom torchvision.models import resnet18\nfrom torchvision.ops import roi_align\n\nT_START = time.time(); MAX_HOURS = 10.5\ndev = \"cuda\"\ntorch.backends.cudnn.benchmark = True\n\nOUT = \"/kaggle/input/notebooks/subashreevjc/imagenetresnet18/shards\"\ntrain_urls = sorted(glob.glob(f\"{OUT}/train-*.tar\"))\nval_urls   = sorted(glob.glob(f\"{OUT}/val-*.tar\"))\nprint(len(train_urls), len(val_urls))   # expect 129 5\n\nN_TRAIN, BS, EPOCHS = 1281167, 256, 5\nsteps_per_epoch = N_TRAIN // BS\nCKPT = \"/kaggle/working/resnet18_ckpt.pth\"\nITER_LOG = \"/kaggle/working/iter_log.csv\"\nPRINT_EVERY = 100         # print every iteration (CSV always has all of them)\nPEAK_LR, WD, MIN_SCALE = 0.2, 5e-5, 0.30\n\ndef res_for_epoch(e):      # 128 -> 160 -> 192, last epochs at 192\n    return [128, 160, 160, 192, 192][min(e, 4)]\n\ndef to_tensor(img):\n    return torch.from_numpy(np.asarray(img).copy()).permute(2, 0, 1)\n\ndef make_loader(urls, train, bs=256, workers=3):\n    ds = wds.WebDataset(urls, shardshuffle=100 if train else False,\n                        handler=wds.warn_and_continue)\n    if train: ds = ds.shuffle(1000)\n    ds = ds.decode(\"pil\").to_tuple(\"jpg\", \"cls\").map_tuple(to_tensor, int)\n    return torch.utils.data.DataLoader(ds, batch_size=bs, num_workers=workers,\n                                       pin_memory=True, drop_last=train,\n                                       prefetch_factor=2, persistent_workers=False)\n\nval_loader = make_loader(val_urls, False, 250, workers=2)\nmean = torch.tensor([0.485,0.456,0.406], device=dev).view(1,3,1,1) * 255\nstd  = torch.tensor([0.229,0.224,0.225], device=dev).view(1,3,1,1) * 255\n\ndef prep(x):               # uint8 -> normalized (used for validation)\n    return ((x.float() - mean) / std).contiguous(memory_format=torch.channels_last)\n\ndef gpu_rrc_flip(x, res):  # per-sample random-resized-crop + flip on the GPU\n    B = x.size(0); x = x.float()\n    area = torch.empty(B, device=dev).uniform_(MIN_SCALE, 1.0)\n    ratio = torch.exp(torch.empty(B, device=dev).uniform_(math.log(3/4), math.log(4/3)))\n    w = (torch.sqrt(area * ratio) * 224).clamp(max=224)\n    h = (torch.sqrt(area / ratio) * 224).clamp(max=224)\n    x1 = torch.rand(B, device=dev) * (224 - w)\n    y1 = torch.rand(B, device=dev) * (224 - h)\n    boxes = torch.stack([torch.arange(B, device=dev, dtype=torch.float32),\n                         x1, y1, x1 + w, y1 + h], 1)\n    out = roi_align(x, boxes, (res, res), spatial_scale=1.0, sampling_ratio=2, aligned=True)\n    flip = torch.rand(B, device=dev) < 0.5\n    out = torch.where(flip.view(-1,1,1,1), out.flip(3), out)\n    return ((out - mean) / std).contiguous(memory_format=torch.channels_last)\n\nmodel = resnet18(weights=None, num_classes=1000, zero_init_residual=True)\nmodel = model.to(dev).to(memory_format=torch.channels_last)\n\ndecay    = [p for p in model.parameters() if p.ndim > 1]\nno_decay = [p for p in model.parameters() if p.ndim <= 1]     # BN + bias\nopt = torch.optim.SGD([{\"params\": decay, \"weight_decay\": WD},\n                       {\"params\": no_decay, \"weight_decay\": 0.0}],\n                      lr=PEAK_LR, momentum=0.9, nesterov=True)\nsched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=PEAK_LR,\n            total_steps=EPOCHS * steps_per_epoch, pct_start=0.17,\n            anneal_strategy=\"linear\", div_factor=20, final_div_factor=1000)\nscaler = torch.amp.GradScaler(\"cuda\")\ncrit = nn.CrossEntropyLoss(label_smoothing=0.1)\n\nstart_epoch = 0\nhistory = {\"loss\": [], \"train_acc\": [], \"val_acc\": [], \"val_top5\": [], \"iter_loss\": []}\nif os.path.exists(CKPT):\n    ck = torch.load(CKPT, map_location=dev)\n    model.load_state_dict(ck[\"model\"]); opt.load_state_dict(ck[\"opt\"])\n    sched.load_state_dict(ck[\"sched\"]); scaler.load_state_dict(ck[\"scaler\"])\n    start_epoch = ck[\"epoch\"]; history.update(ck[\"history\"])\n    print(\"resumed from epoch\", start_epoch)\nif start_epoch == 0:\n    with open(ITER_LOG, \"w\", newline=\"\") as f:\n        csv.writer(f).writerow([\"epoch\", \"iter\", \"res\", \"loss\", \"batch_acc\", \"lr\", \"img_per_s\"])\n\ndef evaluate():            # full 224 image, averaged with its horizontal flip\n    model.eval(); c1 = c5 = seen = 0\n    with torch.no_grad(), torch.autocast(\"cuda\"):\n        for x, y in val_loader:\n            x = prep(x.to(dev, non_blocking=True)); y = y.to(dev)\n            out = model(x).float().softmax(1) + model(x.flip(3)).float().softmax(1)\n            top5 = out.topk(5, 1).indices\n            c1 += (top5[:, 0] == y).sum().item()\n            c5 += (top5 == y[:, None]).any(1).sum().item(); seen += len(y)\n    return c1 / seen, c5 / seen\n\nfor epoch in range(start_epoch, EPOCHS):\n    res = res_for_epoch(epoch)\n    train_loader = make_loader(train_urls, True, BS, workers=3)\n    model.train(); tot_loss = tot_correct = tot_seen = n = 0; t0 = time.time()\n    logf = open(ITER_LOG, \"a\", newline=\"\"); w = csv.writer(logf)\n    print(f\"--- epoch {epoch+1}/{EPOCHS} at {res}px ---\", flush=True)\n    for x, y in train_loader:\n        x = gpu_rrc_flip(x.to(dev, non_blocking=True), res)\n        y = y.to(dev, non_blocking=True)\n        with torch.autocast(\"cuda\"):\n            out = model(x); loss = crit(out, y)\n        opt.zero_grad(set_to_none=True)\n        scaler.scale(loss).backward(); scaler.step(opt); scaler.update()\n        if sched.last_epoch < EPOCHS * steps_per_epoch - 1: sched.step()\n        n += 1; l = loss.item()\n        correct = (out.argmax(1) == y).sum().item()\n        b_acc = correct / len(y)\n        tot_loss += l; tot_correct += correct; tot_seen += len(y)\n        lr = opt.param_groups[0][\"lr\"]; ips = tot_seen / (time.time() - t0)\n        history[\"iter_loss\"].append(l)\n        w.writerow([epoch+1, n, res, f\"{l:.4f}\", f\"{b_acc:.4f}\", f\"{lr:.5f}\", f\"{ips:.0f}\"])\n        if n % PRINT_EVERY == 0:\n            print(f\"  [epoch {epoch+1}] iter {n}/{steps_per_epoch}  loss {l:.3f}  \"\n                  f\"avg {tot_loss/n:.3f}  acc {b_acc:.3f}  lr {lr:.4f}  {ips:.0f} img/s\", flush=True)\n    logf.close()\n    val1, val5 = evaluate()\n    history[\"loss\"].append(tot_loss/n); history[\"train_acc\"].append(tot_correct/tot_seen)\n    history[\"val_acc\"].append(val1); history[\"val_top5\"].append(val5)\n    print(f\"\\n=== EPOCH {epoch+1}/{EPOCHS} | {res}px | train loss {tot_loss/n:.3f} | \"\n          f\"train acc {tot_correct/tot_seen:.3f} | val top-1 {val1:.3f} | val top-5 {val5:.3f} | \"\n          f\"{(time.time()-t0)/60:.1f} min ===\\n\", flush=True)\n    torch.save({\"model\": model.state_dict(), \"opt\": opt.state_dict(),\n                \"sched\": sched.state_dict(), \"scaler\": scaler.state_dict(),\n                \"epoch\": epoch+1, \"history\": history}, CKPT)\n    del train_loader; gc.collect()\n    if (time.time() - T_START) / 3600 > MAX_HOURS:\n        print(\"Time limit near, stopping.\"); break\n\ntorch.save(model.state_dict(), \"/kaggle/working/resnet18_final.pth\")\nfig, ax = plt.subplots(1, 3, figsize=(16, 4))\nax[0].plot(history[\"iter_loss\"], lw=0.5); ax[0].set_title(\"Loss per iteration\")\nax[1].plot(history[\"loss\"], marker=\"o\"); ax[1].set_title(\"Loss per epoch\")\nax[2].plot(history[\"train_acc\"], marker=\"o\", label=\"train\"); ax[2].plot(history[\"val_acc\"], marker=\"o\", label=\"val top-1\")\nax[2].plot(history[\"val_top5\"], marker=\"o\", label=\"val top-5\"); ax[2].set_title(\"Accuracy\"); ax[2].legend()\nplt.savefig(\"/kaggle/working/curves.png\", dpi=120); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T10:32:42.116541Z","iopub.execute_input":"2026-09-29T10:32:42.11699Z","iopub.status.idle":"2026-09-29T11:42:22.306219Z","shell.execute_reply.started":"2026-09-29T10:32:42.116962Z","shell.execute_reply":"2026-09-29T11:42:22.305447Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T11:43:52.176892Z","iopub.status.idle":"2026-09-29T11:43:52.177258Z","shell.execute_reply.started":"2026-09-29T11:43:52.177121Z","shell.execute_reply":"2026-09-29T11:43:52.17714Z"}},"outputs":[],"execution_count":null}]}