{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"! pip install -q timm \\\n  https://pip.repos.neuron.amazonaws.com/torch-xla/torch_xla-1.13.0%2Btorchneuron3-cp37-cp37m-linux_x86_64.whl","metadata":{"_uuid":"e26026ac-4bb0-4712-a9f3-5adf2d5786cf","_cell_guid":"2df78fc0-1c9c-4518-89f2-e7db72f3ecc0","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-03-18T20:21:19.986392Z","iopub.execute_input":"2023-03-18T20:21:19.987317Z","iopub.status.idle":"2023-03-18T20:21:46.123187Z","shell.execute_reply.started":"2023-03-18T20:21:19.987276Z","shell.execute_reply":"2023-03-18T20:21:46.121536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! export XLA_USE_BF16=1","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import glob\nimport io\n\nimport pandas as pd\n\nimport tensorflow as tf\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nimport torch_xla.core.xla_model as xm\nimport torchvision.transforms as T\n\nfrom joblib import Parallel, delayed\nfrom PIL import Image\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import StratifiedKFold\nfrom tqdm.auto import tqdm\nfrom sklearn.metrics import f1_score","metadata":{"_uuid":"82a56d80-9525-49be-8ba3-1f44e31f56ec","_cell_guid":"f127c779-06ab-4f60-a6ad-b1c31fcff06e","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-03-18T20:21:46.126857Z","iopub.execute_input":"2023-03-18T20:21:46.127657Z","iopub.status.idle":"2023-03-18T20:21:58.823456Z","shell.execute_reply.started":"2023-03-18T20:21:46.127617Z","shell.execute_reply":"2023-03-18T20:21:58.822203Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class config:\n    MAX_EPOCHS = 5      \n    N_SPLITS = 5          # Must be equal or less than 8\n    BATCH_SIZE = 32","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def tfrecords_to_dataframe(fp, test=False):\n\n    def parse(pb, test=False):\n        d = {\n            \"id\": tf.io.FixedLenFeature([], tf.string),\n            \"image\": tf.io.FixedLenFeature([], tf.string),\n        }\n        if not test:\n            d[\"class\"] = tf.io.FixedLenFeature([], tf.int64)\n        return tf.io.parse_single_example(pb, d)\n\n    df = {\"id\": [], \"img\": []}\n    if not test:\n        df[\"lab\"] = []\n    for sample in tf.data.TFRecordDataset(glob.glob(fp)).map(\n        lambda pb: parse(pb, test)\n    ):\n        df[\"id\"].append(sample[\"id\"].numpy().decode(\"utf-8\"))\n        df[\"img\"].append(sample[\"image\"].numpy())\n        if not test:\n            df[\"lab\"].append(sample[\"class\"].numpy())\n    return pd.DataFrame(df)","metadata":{"_uuid":"8f5f25cb-a27c-44a6-9baf-057ea0cb1b03","_cell_guid":"df704c0c-6187-4211-a673-cfd624153cde","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-03-18T20:21:58.825149Z","iopub.execute_input":"2023-03-18T20:21:58.825563Z","iopub.status.idle":"2023-03-18T20:21:58.835966Z","shell.execute_reply.started":"2023-03-18T20:21:58.825521Z","shell.execute_reply":"2023-03-18T20:21:58.834889Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AverageMeter():\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count","metadata":{"execution":{"iopub.status.busy":"2023-03-18T20:21:58.841112Z","iopub.execute_input":"2023-03-18T20:21:58.841733Z","iopub.status.idle":"2023-03-18T20:21:58.849779Z","shell.execute_reply.started":"2023-03-18T20:21:58.841693Z","shell.execute_reply":"2023-03-18T20:21:58.848567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.concat([\n    tfrecords_to_dataframe(\"../input/tpu-getting-started/tfrecords-jpeg-224x224/train/*.tfrec\"),\n    tfrecords_to_dataframe(\"../input/tpu-getting-started/tfrecords-jpeg-224x224/val/*.tfrec\"),\n], ignore_index=True).reset_index()\n\ntrain_df.drop('index', axis=1, inplace=True)\n\ncv = StratifiedKFold(n_splits=config.N_SPLITS, random_state=42, shuffle=True)\n\ntrain_df['fold'] = -1\n\nfor fold, (train_idx, val_idx) in enumerate(cv.split(train_df, train_df['lab'])):\n    train_df.loc[val_idx, 'fold'] = fold","metadata":{"_uuid":"2836afb4-9780-4ac6-a322-b32b840b2327","_cell_guid":"4063030d-a384-4265-a959-d0b846391f82","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-03-18T20:21:58.851441Z","iopub.execute_input":"2023-03-18T20:21:58.852664Z","iopub.status.idle":"2023-03-18T20:22:10.424946Z","shell.execute_reply.started":"2023-03-18T20:21:58.852595Z","shell.execute_reply":"2023-03-18T20:22:10.423688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PetalDataset(Dataset):\n    \n    def __init__(self, df, test=False):\n        self.df = df\n        self.test = test\n        self.transform = T.ToTensor()\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        item = self.df.iloc[idx]\n        img = Image.open(io.BytesIO(item.img))\n        img = self.transform(img)\n        \n        if self.test:\n            return img\n        \n        label = item.lab\n        return img, label","metadata":{"_uuid":"ceafe78d-aa01-434f-9c0b-7619c5322cbf","_cell_guid":"5f5579c4-5dae-44f9-ba48-297b64dee4ae","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-03-18T20:22:10.426575Z","iopub.execute_input":"2023-03-18T20:22:10.426982Z","iopub.status.idle":"2023-03-18T20:22:10.435034Z","shell.execute_reply.started":"2023-03-18T20:22:10.426945Z","shell.execute_reply":"2023-03-18T20:22:10.433907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Downloading the model\nmodel = timm.create_model('efficientnet_b0', pretrained=True)\ndel model","metadata":{"_uuid":"cf45abef-f9b9-45d2-b816-1580e3026a08","_cell_guid":"43c5bb51-8f9e-434e-a499-1c377ac859d8","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-03-18T20:22:10.437442Z","iopub.execute_input":"2023-03-18T20:22:10.438190Z","iopub.status.idle":"2023-03-18T20:22:11.304573Z","shell.execute_reply.started":"2023-03-18T20:22:10.438153Z","shell.execute_reply":"2023-03-18T20:22:11.303128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"counts = train_df['lab'].value_counts()\nclass_weights = torch.tensor(1/counts.sort_index(), dtype=torch.float)","metadata":{"execution":{"iopub.status.busy":"2023-03-18T20:22:11.306273Z","iopub.execute_input":"2023-03-18T20:22:11.306727Z","iopub.status.idle":"2023-03-18T20:22:11.320172Z","shell.execute_reply.started":"2023-03-18T20:22:11.306671Z","shell.execute_reply":"2023-03-18T20:22:11.318763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_fn(fold):\n    device = xm.xla_device(fold + 1)\n\n    val_ = train_df.query('fold == @fold')\n    train_ = train_df.query('fold != @fold')\n\n    train_ds = PetalDataset(train_)\n    val_ds = PetalDataset(val_)\n\n    train_loader = DataLoader(train_ds, batch_size=config.BATCH_SIZE, drop_last=True, shuffle=True)\n    val_loader = DataLoader(val_ds, batch_size=config.BATCH_SIZE)\n\n    model = timm.create_model('efficientnet_b0', pretrained=True, num_classes=104)\n    model.to(device)\n    optimizer = optim.Adam(model.parameters())\n    scheduler = optim.lr_scheduler.OneCycleLR(\n        optimizer,\n        max_lr=1e-4,\n        steps_per_epoch=len(train_loader),\n        pct_start=0.1,\n        epochs=config.MAX_EPOCHS,\n    )\n    loss_fn = nn.CrossEntropyLoss(weight=class_weights)\n    meter = AverageMeter()\n    \n    for epoch in range(config.MAX_EPOCHS):\n        model.train()\n        stream = tqdm(train_loader, desc=f\"Fold={fold}, Epoch={epoch}\")\n        for data, target in stream:\n            optimizer.zero_grad()\n            data = data.to(device)\n            target = target.to(device)\n            output = model(data)\n            loss = loss_fn(output, target)\n            loss.backward()\n\n            optimizer.step()\n            xm.mark_step()\n            scheduler.step()\n            \n            meter.update(loss.item())\n            \n            stream.set_postfix({\n                \"train_loss\": meter.avg,\n            })\n        meter.reset()\n        \n        model.eval()\n        with torch.no_grad():\n            stream = tqdm(val_loader, desc=f\"Fold={fold}, Validating...\")\n            for data, target in stream:\n                data = data.to(device)\n                target = target.to(device)\n                output = model(data)\n                loss = loss_fn(output, target)\n                meter.update(loss.item())\n\n                stream.set_postfix({\n                    \"val_loss\": meter.avg,\n                })\n        meter.reset()\n            \n    model.eval()\n    xm.save(model.state_dict(), f\"model_fold_{fold}.pt\")\n        \n    y_true = []\n    y_pred = []\n\n    with torch.no_grad():\n        for data, target in stream:\n            data = data.to(device)\n            target = target.to(device)\n            output = model(data)\n            y_pred.extend(output.argmax(axis=1).cpu().numpy())\n            y_true.extend(target.squeeze().cpu().numpy())\n\n    val_score = f1_score(y_true, y_pred, average='macro')      \n    return val_score","metadata":{"_uuid":"efa66142-1068-460d-8d57-139ddca3d2a1","_cell_guid":"13341f9b-75dd-4737-b0ad-f7fc9d3826d7","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-03-18T20:22:11.322074Z","iopub.execute_input":"2023-03-18T20:22:11.322803Z","iopub.status.idle":"2023-03-18T20:22:11.336309Z","shell.execute_reply.started":"2023-03-18T20:22:11.322760Z","shell.execute_reply":"2023-03-18T20:22:11.335162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"Parallel(n_jobs=config.N_SPLITS, backend=\"threading\")(delayed(train_fn)(i) for i in range(config.N_SPLITS))","metadata":{"_uuid":"ce55fd3d-5dd5-4999-9f97-d391ec3ec18f","_cell_guid":"ab1dc678-70d5-4bfb-8ce0-a50ae8fa68bd","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-03-18T20:22:11.338898Z","iopub.execute_input":"2023-03-18T20:22:11.339561Z"},"trusted":true},"execution_count":null,"outputs":[]}]}