{"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":"gpu","dataSources":[{"sourceType":"competition","sourceId":21154,"databundleVersionId":1243559}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-04-13T14:51:59.741542Z","iopub.execute_input":"2026-04-13T14:51:59.741860Z","iopub.status.idle":"2026-04-13T14:52:01.610616Z","shell.execute_reply.started":"2026-04-13T14:51:59.741830Z","shell.execute_reply":"2026-04-13T14:52:01.609813Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport tensorflow as tf\n\nprint(\"TF version:\", tf.__version__)\nprint(\"Logical TPU devices:\", tf.config.list_logical_devices(\"TPU\"))\nprint(\"All logical devices:\", tf.config.list_logical_devices())\n\ntry:\n    resolver = tf.distribute.cluster_resolver.TPUClusterResolver()\n    print(\"TPU master:\", resolver.master())\n    tf.config.experimental_connect_to_cluster(resolver)\n    tf.tpu.experimental.initialize_tpu_system(resolver)\n    print(\"Initialized TPU devices:\", tf.config.list_logical_devices(\"TPU\"))\n    strategy = tf.distribute.TPUStrategy(resolver)\nexcept Exception as e:\n    print(\"TPU init failed:\", repr(e))\n    strategy = tf.distribute.get_strategy()\n\nprint(\"REPLICAS:\", strategy.num_replicas_in_sync)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-13T14:52:01.612323Z","iopub.execute_input":"2026-04-13T14:52:01.612628Z","iopub.status.idle":"2026-04-13T14:52:40.428355Z","shell.execute_reply.started":"2026-04-13T14:52:01.612607Z","shell.execute_reply":"2026-04-13T14:52:40.427594Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# install some more useful packages\n\n!pip install -q timm\n\nimport os\nimport io\nimport glob\nimport shutil\nfrom pathlib import Path\n\nimport tensorflow as tf\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader\nfrom torchvision import datasets, transforms\n\nimport timm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-13T14:52:40.429215Z","iopub.execute_input":"2026-04-13T14:52:40.429722Z","iopub.status.idle":"2026-04-13T14:52:58.156208Z","shell.execute_reply.started":"2026-04-13T14:52:40.429699Z","shell.execute_reply":"2026-04-13T14:52:58.155341Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# use the medium size pictures\n\nDATA_DIR = \"/kaggle/input/competitions/tpu-getting-started/tfrecords-jpeg-224x224\"\nWORK_DIR = Path(\"/kaggle/working/petals_images\")\nWORK_DIR.mkdir(parents=True, exist_ok=True)\n\nprint(os.listdir(DATA_DIR))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-13T14:52:58.157257Z","iopub.execute_input":"2026-04-13T14:52:58.157834Z","iopub.status.idle":"2026-04-13T14:52:58.163220Z","shell.execute_reply.started":"2026-04-13T14:52:58.157806Z","shell.execute_reply":"2026-04-13T14:52:58.162577Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# swin wants jpeg format pictures, so convert tfrecs to jpegs\n\ndef get_feature(example, names):\n    for name in names:\n        if name in example.features.feature:\n            f = example.features.feature[name]\n            if f.bytes_list.value:\n                return f.bytes_list.value[0]\n            if f.int64_list.value:\n                return int(f.int64_list.value[0])\n            if f.float_list.value:\n                return float(f.float_list.value[0])\n    return None\n\ndef export_tfrecords_to_images(tfrecord_paths, out_root, labeled=True):\n    out_root = Path(out_root)\n    out_root.mkdir(parents=True, exist_ok=True)\n\n    count = 0\n    for tfrecord_path in tfrecord_paths:\n        ds = tf.data.TFRecordDataset(tfrecord_path)\n        for raw_record in ds:\n            ex = tf.train.Example()\n            ex.ParseFromString(raw_record.numpy())\n\n            image_bytes = get_feature(ex, [\"img\", \"image\"])\n            image_id = get_feature(ex, [\"id\"])\n            label = get_feature(ex, [\"label\", \"class\"]) if labeled else None\n\n            if image_bytes is None:\n                continue\n\n            if isinstance(image_id, bytes):\n                image_id = image_id.decode(\"utf-8\")\n\n            if image_id is None:\n                image_id = f\"sample_{count:06d}\"\n\n            if labeled:\n                class_dir = out_root / str(label)\n                class_dir.mkdir(parents=True, exist_ok=True)\n                out_path = class_dir / f\"{image_id}.jpg\"\n            else:\n                out_path = out_root / f\"{image_id}.jpg\"\n\n            with open(out_path, \"wb\") as f:\n                f.write(image_bytes)\n\n            count += 1\n\n    print(f\"Exported {count} images to {out_root}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-13T14:52:58.164114Z","iopub.execute_input":"2026-04-13T14:52:58.164438Z","iopub.status.idle":"2026-04-13T14:52:58.220145Z","shell.execute_reply.started":"2026-04-13T14:52:58.164403Z","shell.execute_reply":"2026-04-13T14:52:58.219448Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_files = sorted(glob.glob(os.path.join(DATA_DIR, \"train/*.tfrec\"))) + \\\n              sorted(glob.glob(os.path.join(DATA_DIR, \"train/*.tfrecord\")))\n\nval_files = sorted(glob.glob(os.path.join(DATA_DIR, \"val/*.tfrec\"))) + \\\n            sorted(glob.glob(os.path.join(DATA_DIR, \"val/*.tfrecord\")))\n\ntest_files = sorted(glob.glob(os.path.join(DATA_DIR, \"test/*.tfrec\"))) + \\\n             sorted(glob.glob(os.path.join(DATA_DIR, \"test/*.tfrecord\")))\n\nprint(os.path.join(DATA_DIR, \"train/\"))\nprint(\"train:\", len(train_files))\nprint(\"val:\", len(val_files))\nprint(\"test:\", len(test_files))\n\nexport_tfrecords_to_images(train_files, WORK_DIR / \"train\", labeled=True)\nexport_tfrecords_to_images(val_files, WORK_DIR / \"val\", labeled=True)\nexport_tfrecords_to_images(test_files, WORK_DIR / \"test\", labeled=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-13T14:52:58.221034Z","iopub.execute_input":"2026-04-13T14:52:58.221217Z","iopub.status.idle":"2026-04-13T14:53:09.377974Z","shell.execute_reply.started":"2026-04-13T14:52:58.221199Z","shell.execute_reply":"2026-04-13T14:53:09.377296Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IMG_SIZE = 224\nBATCH_SIZE = 32\nNUM_CLASSES = 104\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\ntrain_tfms = transforms.Compose([\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.RandomHorizontalFlip(),\n    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.05),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=(0.485, 0.456, 0.406),\n                         std=(0.229, 0.224, 0.225)),\n])\n\nval_tfms = transforms.Compose([\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=(0.485, 0.456, 0.406),\n                         std=(0.229, 0.224, 0.225)),\n])\n\ntrain_ds = datasets.ImageFolder(WORK_DIR / \"train\", transform=train_tfms)\nval_ds   = datasets.ImageFolder(WORK_DIR / \"val\", transform=val_tfms)\n\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True, num_workers=2, pin_memory=True)\nval_loader   = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False, num_workers=2, pin_memory=True)\n\nprint(\"train classes:\", len(train_ds.classes))\nprint(\"val classes:\", len(val_ds.classes))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-13T14:53:09.379704Z","iopub.execute_input":"2026-04-13T14:53:09.380026Z","iopub.status.idle":"2026-04-13T14:53:09.428588Z","shell.execute_reply.started":"2026-04-13T14:53:09.380002Z","shell.execute_reply":"2026-04-13T14:53:09.428043Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = timm.create_model(\n    \"swin_tiny_patch4_window7_224\",\n    pretrained=True,\n    num_classes=NUM_CLASSES\n).to(DEVICE)\n\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4)\nscaler = torch.cuda.amp.GradScaler(enabled=torch.cuda.is_available())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-13T14:53:09.429390Z","iopub.execute_input":"2026-04-13T14:53:09.429582Z","iopub.status.idle":"2026-04-13T14:53:13.762100Z","shell.execute_reply.started":"2026-04-13T14:53:09.429563Z","shell.execute_reply":"2026-04-13T14:53:13.761434Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, loader):\n    model.train()\n    running_loss, running_correct, total = 0.0, 0, 0\n\n    for x, y in loader:\n        x = x.to(DEVICE, non_blocking=True)\n        y = y.to(DEVICE, non_blocking=True)\n\n        optimizer.zero_grad()\n\n        with torch.cuda.amp.autocast(enabled=torch.cuda.is_available()):\n            logits = model(x)\n            loss = criterion(logits, y)\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        running_loss += loss.item() * x.size(0)\n        running_correct += (logits.argmax(1) == y).sum().item()\n        total += x.size(0)\n\n    return running_loss / total, running_correct / total\n\n@torch.no_grad()\ndef evaluate(model, loader):\n    model.eval()\n    running_loss, running_correct, total = 0.0, 0, 0\n\n    for x, y in loader:\n        x = x.to(DEVICE, non_blocking=True)\n        y = y.to(DEVICE, non_blocking=True)\n\n        logits = model(x)\n        loss = criterion(logits, y)\n\n        running_loss += loss.item() * x.size(0)\n        running_correct += (logits.argmax(1) == y).sum().item()\n        total += x.size(0)\n\n    return running_loss / total, running_correct / total\n\nEPOCHS = 16\n\nfor epoch in range(EPOCHS):\n    tr_loss, tr_acc = train_one_epoch(model, train_loader)\n    va_loss, va_acc = evaluate(model, val_loader)\n\n    print(\n        f\"Epoch {epoch+1}/{EPOCHS} | \"\n        f\"train_loss={tr_loss:.4f} train_acc={tr_acc:.4f} | \"\n        f\"val_loss={va_loss:.4f} val_acc={va_acc:.4f}\"\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-13T15:13:08.721285Z","iopub.execute_input":"2026-04-13T15:13:08.721939Z","iopub.status.idle":"2026-04-13T15:42:08.684341Z","shell.execute_reply.started":"2026-04-13T15:13:08.721903Z","shell.execute_reply":"2026-04-13T15:42:08.683433Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from PIL import Image\n\ntest_dir = WORK_DIR / \"test\"\ntest_files = sorted(test_dir.glob(\"*.jpg\"))\n\ninfer_tfms = val_tfms\n\n@torch.no_grad()\ndef predict_one(path):\n    img = Image.open(path).convert(\"RGB\")\n    x = infer_tfms(img).unsqueeze(0).to(DEVICE)\n    logits = model(x)\n    pred = logits.argmax(1).item()\n    return pred\n\nrows = []\nfor path in test_files:\n    image_id = path.stem\n    pred = predict_one(path)\n    rows.append((image_id, pred))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-13T15:51:37.100792Z","iopub.execute_input":"2026-04-13T15:51:37.101457Z","iopub.status.idle":"2026-04-13T15:53:11.170902Z","shell.execute_reply.started":"2026-04-13T15:51:37.101424Z","shell.execute_reply":"2026-04-13T15:53:11.170100Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub = pd.DataFrame(rows, columns=[\"id\", \"label\"])\nsub.to_csv(\"/kaggle/working/submission.csv\", index=False)\nsub.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-13T15:53:19.294090Z","iopub.execute_input":"2026-04-13T15:53:19.294845Z","iopub.status.idle":"2026-04-13T15:53:19.356555Z","shell.execute_reply.started":"2026-04-13T15:53:19.294816Z","shell.execute_reply":"2026-04-13T15:53:19.355941Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}