{"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":"markdown","source":"# 1.Import","metadata":{}},{"cell_type":"code","source":"import os\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'\n\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision\nfrom torchvision import models, transforms\n\nimport tensorflow as tf\ntf.config.experimental.set_visible_devices([], 'GPU')\n\nimport copy\nimport glob\nimport numpy as np\nfrom PIL import Image\nimport io\nimport matplotlib.pyplot as plt\nimport pandas as pd","metadata":{"execution":{"iopub.status.busy":"2023-06-26T19:36:39.340223Z","iopub.execute_input":"2023-06-26T19:36:39.340710Z","iopub.status.idle":"2023-06-26T19:36:39.348750Z","shell.execute_reply.started":"2023-06-26T19:36:39.340678Z","shell.execute_reply":"2023-06-26T19:36:39.347071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2.Create dataset","metadata":{}},{"cell_type":"code","source":"train_files = glob.glob('/kaggle/input/tpu-getting-started/*/train/*.tfrec')\nval_files = glob.glob('/kaggle/input/tpu-getting-started/*/val/*.tfrec')\ntest_files = glob.glob('/kaggle/input/tpu-getting-started/*/test/*.tfrec')","metadata":{"execution":{"iopub.status.busy":"2023-06-26T17:59:16.990606Z","iopub.execute_input":"2023-06-26T17:59:16.991674Z","iopub.status.idle":"2023-06-26T17:59:17.184356Z","shell.execute_reply.started":"2023-06-26T17:59:16.991639Z","shell.execute_reply":"2023-06-26T17:59:17.183481Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Dataset(Dataset):\n    def __init__(self, filenames, transform=None, is_labeled=True):\n        self.filenames = filenames\n        self.transform = transform\n        self.is_labeled = is_labeled\n        self.records = self.parse_files()\n    \n    def parse_files(self):\n        train_feature_description = {\n            'class': tf.io.FixedLenFeature([], tf.int64),\n            'id': tf.io.FixedLenFeature([], tf.string),\n            'image': tf.io.FixedLenFeature([], tf.string)\n        }\n        test_feature_description = {\n            'id': tf.io.FixedLenFeature([], tf.string),\n            'image': tf.io.FixedLenFeature([], tf.string)\n        }\n        records = []\n        for filename in self.filenames:\n            file = tf.data.TFRecordDataset(filename)\n            if self.is_labeled:\n                file = list(file.map(lambda x: tf.io.parse_single_example(x, train_feature_description)))\n            else:\n                file = list(file.map(lambda x: tf.io.parse_single_example(x, test_feature_description)))\n            records += file\n        return records\n    \n    def __len__(self):\n        return len(self.records)\n    \n    def __getitem__(self, index):\n        record = self.records[index]\n        if self.transform:\n            x = self.transform(Image.open(io.BytesIO(record['image'].numpy())))\n        else:\n            x = Image.open(io.BytesIO(record['image'].numpy()))\n        if self.is_labeled:\n            return x, record['class'].numpy()\n        else:\n            return x, record['id'].numpy().decode('utf-8')","metadata":{"execution":{"iopub.status.busy":"2023-06-26T17:59:17.187782Z","iopub.execute_input":"2023-06-26T17:59:17.189815Z","iopub.status.idle":"2023-06-26T17:59:17.202231Z","shell.execute_reply.started":"2023-06-26T17:59:17.189790Z","shell.execute_reply":"2023-06-26T17:59:17.201363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def add_noise(x):\n    noise = torch.randn_like(x)*0.1\n    return x + noise\n\ntrain_transform = transforms.Compose([\n    transforms.Resize((299, 299)),\n    transforms.RandomHorizontalFlip(),\n    transforms.ToTensor(),\n    transforms.Lambda(add_noise),\n    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n])\ntest_transform = transforms.Compose([\n    transforms.Resize((299, 299)),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n])","metadata":{"execution":{"iopub.status.busy":"2023-06-26T17:59:17.206330Z","iopub.execute_input":"2023-06-26T17:59:17.206627Z","iopub.status.idle":"2023-06-26T17:59:17.215308Z","shell.execute_reply.started":"2023-06-26T17:59:17.206604Z","shell.execute_reply":"2023-06-26T17:59:17.214425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = Dataset(train_files, train_transform)\nval_dataset = Dataset(val_files, test_transform)","metadata":{"execution":{"iopub.status.busy":"2023-06-26T17:59:17.219082Z","iopub.execute_input":"2023-06-26T17:59:17.219524Z","iopub.status.idle":"2023-06-26T18:00:37.505436Z","shell.execute_reply.started":"2023-06-26T17:59:17.219501Z","shell.execute_reply":"2023-06-26T18:00:37.504430Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataloader = DataLoader(train_dataset, batch_size=66, shuffle=True)\nval_dataloader = DataLoader(val_dataset, batch_size=66, shuffle=True)\nprint(f'Train dataset has {len(train_dataloader.dataset)} images')\nprint(f'Validation dataset has {len(val_dataloader.dataset)} images')","metadata":{"execution":{"iopub.status.busy":"2023-06-26T18:00:37.508650Z","iopub.execute_input":"2023-06-26T18:00:37.509538Z","iopub.status.idle":"2023-06-26T18:00:37.516292Z","shell.execute_reply.started":"2023-06-26T18:00:37.509503Z","shell.execute_reply":"2023-06-26T18:00:37.515083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataloader_iter = iter(train_dataloader)\nx, y = next(train_dataloader_iter)\n\nplt.figure(figsize=(20,30))\nplt.imshow(np.transpose(torchvision.utils.make_grid(x, nrow=11), (1, 2, 0)));","metadata":{"execution":{"iopub.status.busy":"2023-06-26T18:00:37.517870Z","iopub.execute_input":"2023-06-26T18:00:37.518871Z","iopub.status.idle":"2023-06-26T18:00:40.875348Z","shell.execute_reply.started":"2023-06-26T18:00:37.518838Z","shell.execute_reply":"2023-06-26T18:00:40.874113Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(f\"Using {device} device\")","metadata":{"execution":{"iopub.status.busy":"2023-06-26T18:00:40.876710Z","iopub.execute_input":"2023-06-26T18:00:40.877097Z","iopub.status.idle":"2023-06-26T18:00:40.888668Z","shell.execute_reply.started":"2023-06-26T18:00:40.877061Z","shell.execute_reply":"2023-06-26T18:00:40.887804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3.Create Model","metadata":{}},{"cell_type":"code","source":"class EarlyStopper:\n    def __init__(self, model, patience, threshold, is_abs=False):\n        self.patience = patience\n        self.threshold = threshold\n        self.counter = 0\n        self.best_loss = 1\n        self.model = model\n        self.state_dict = copy.deepcopy(model.state_dict())\n        self.stopped = False\n        self.is_abs=is_abs\n\n    def check(self, loss):\n        if not self.stopped:\n            if not self.is_abs:\n                difference = self.best_loss * (1 - self.threshold)\n            else:\n                difference = self.best_loss - self.threshold\n\n            if loss < difference:\n                self.best_loss = loss\n                self.state_dict = copy.deepcopy(self.model.state_dict())\n                self.counter = 0\n            else:\n                if self.counter < self.patience-1:\n                    self.counter += 1\n                else:\n                    self.stopped = True\n        return self.stopped\n    \n    def get_state_dict(self):\n        return self.state_dict\n    \n    def zero_counter(self):\n        self.counter = 0\n        self.stopped = False\n        self.best_loss = 1","metadata":{"execution":{"iopub.status.busy":"2023-06-26T18:00:40.890347Z","iopub.execute_input":"2023-06-26T18:00:40.891039Z","iopub.status.idle":"2023-06-26T18:00:40.906282Z","shell.execute_reply.started":"2023-06-26T18:00:40.891000Z","shell.execute_reply":"2023-06-26T18:00:40.905373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Model(torch.nn.Module):\n    def __init__(self, n_out):\n        super().__init__()\n        state_dict = torch.hub.load_state_dict_from_url(\n            'https://download.pytorch.org/models/inception_v3_google-1a9a5a14.pth',\n            model_dir='.',\n            check_hash=True\n        )\n        self.model = models.inception_v3()\n        self.model.load_state_dict(state_dict)\n        self.n_out = n_out\n        self.model.fc = torch.nn.Linear(2048, n_out)\n        self.freeze()\n\n    def forward(self, x):\n        x = self.model(x)\n        return x.logits\n        \n    def freeze(self):\n        for name, param in self.model.named_parameters():\n            if not name.startswith(\"fc\"):\n                param.requires_grad = False\n\n    def unfreeze(self, start_index, end_index):\n        for idx, layer in enumerate(self.model.children()):\n            if start_index <= idx <= end_index:\n                for param in layer.parameters():\n                    param.requires_grad = True\n    \n    def train_step(self, train_dataloader, optimizer, criterion):\n        train_loss = 0\n        train_correct = 0\n\n        for x, y in train_dataloader:\n            optimizer.zero_grad()\n            x_train, y_train = x.to(device).float(), torch.nn.functional.one_hot(y, self.n_out).to(device).float()\n            output = self(x_train)\n            loss = criterion(output, y_train)\n            train_loss += loss.item()\n            train_correct += (y.to(device).float() == torch.argmax(torch.nn.functional.softmax(output, dim=1), dim=1)).float().sum()\n            loss.backward()\n            optimizer.step()\n        train_loss = train_loss/len(train_dataloader.dataset)\n        train_correct = train_correct/len(train_dataloader.dataset)\n        return train_loss, train_correct\n    \n    def val_step(self, val_dataloader, criterion):\n        val_loss = 0\n        val_correct = 0\n        \n        with torch.no_grad():\n            for x, y in val_dataloader:\n                x_val, y_val = x.to(device).float(), torch.nn.functional.one_hot(y, self.n_out).to(device).float()\n                pred = self(x_val)\n                val_loss += criterion(pred, y_val).item()\n                val_correct += (y.to(device).float() == torch.argmax(torch.nn.functional.softmax(pred, dim=1), dim=1)).float().sum()\n\n        val_loss = val_loss/len(val_dataloader.dataset)\n        val_correct = val_correct/len(val_dataloader.dataset)\n        return val_loss, val_correct\n    \n    def fit(self, n_epochs, train_dataloader, val_dataloader, optimizer, criterion, start_epoch=0, scheduler=None, early_stopper=None):\n        pb_length = 50\n        lr = optimizer.param_groups[0][\"lr\"]\n        \n        train_loss_plot = []\n        val_loss_plot = []\n        train_accuracy_plot = []\n        val_accuracy_plot = []\n        \n        for epoch in range(start_epoch, start_epoch + n_epochs):\n            train_loss, train_acc = self.train_step(train_dataloader, optimizer, criterion)\n            val_loss, val_acc = self.val_step(val_dataloader, criterion)\n            \n            train_loss_plot.append(train_loss)\n            val_loss_plot.append(val_loss)\n            train_accuracy_plot.append(train_acc)\n            val_accuracy_plot.append(val_acc)\n            \n            if early_stopper:\n                if early_stopper.check(val_loss):\n                    self.load_state_dict(early_stopper.get_state_dict())\n                    break\n            if scheduler:\n                scheduler.step(val_loss)\n                lr = optimizer.param_groups[0][\"lr\"]\n            pb_progress = epoch - start_epoch+ 1\n            pb_percent = pb_length * (pb_progress / n_epochs)\n            pb_bar = \"❚\" * int(pb_percent) + \" \" * (pb_length - int(pb_percent))\n            print(f\"|{pb_bar}| {pb_progress} / {n_epochs}, train_loss = {train_loss:.4f}, train_accuracy = {train_acc:.4f}, val_loss = {val_loss:.4f}, val_accuracy = {val_acc:.4f}, lr = {lr:.3}\", end=\"\\r\")\n        early_stopper.zero_counter()\n        \n        fig, (ax1, ax2) = plt.subplots(2, 1)\n        ax1.plot([i for i in range(start_epoch, start_epoch + n_epochs)], train_loss_plot, label=\"train\")\n        ax1.plot([i for i in range(start_epoch, start_epoch + n_epochs)], val_loss_plot, label=\"val\")\n        ax1.legend()\n        ax2.plot([i for i in range(start_epoch, start_epoch + n_epochs)], train_accuracy_plot, label=\"train\")\n        ax2.plot([i for i in range(start_epoch, start_epoch + n_epochs)], val_accuracy_plot, label=\"val\")\n        ax2.legend()\n        plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-06-26T18:00:40.910699Z","iopub.execute_input":"2023-06-26T18:00:40.911555Z","iopub.status.idle":"2023-06-26T18:00:40.945460Z","shell.execute_reply.started":"2023-06-26T18:00:40.911519Z","shell.execute_reply":"2023-06-26T18:00:40.944342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = Model(104).to(device)","metadata":{"execution":{"iopub.status.busy":"2023-06-26T18:00:40.947183Z","iopub.execute_input":"2023-06-26T18:00:40.948013Z","iopub.status.idle":"2023-06-26T18:00:54.824978Z","shell.execute_reply.started":"2023-06-26T18:00:40.947979Z","shell.execute_reply":"2023-06-26T18:00:54.823954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 4.Train","metadata":{}},{"cell_type":"code","source":"criterion = torch.nn.CrossEntropyLoss()\noptimizer = torch.optim.SGD(model.parameters(), lr=0.03)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', 0.3, 3, threshold=0.01)\nearly_stopper = EarlyStopper(model, patience=8, threshold=0.01)","metadata":{"execution":{"iopub.status.busy":"2023-06-26T18:00:54.826443Z","iopub.execute_input":"2023-06-26T18:00:54.826793Z","iopub.status.idle":"2023-06-26T18:00:54.883713Z","shell.execute_reply.started":"2023-06-26T18:00:54.826760Z","shell.execute_reply":"2023-06-26T18:00:54.882815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.fit(50, train_dataloader, val_dataloader, optimizer, criterion, scheduler=scheduler, early_stopper=early_stopper)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataloader = DataLoader(train_dataset, batch_size=36, shuffle=True)\nval_dataloader = DataLoader(val_dataset, batch_size=36, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2023-06-26T17:57:58.263025Z","iopub.status.idle":"2023-06-26T17:57:58.263544Z","shell.execute_reply.started":"2023-06-26T17:57:58.263273Z","shell.execute_reply":"2023-06-26T17:57:58.263315Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.unfreeze(0, 21)","metadata":{"execution":{"iopub.status.busy":"2023-06-26T17:57:58.265344Z","iopub.status.idle":"2023-06-26T17:57:58.265843Z","shell.execute_reply.started":"2023-06-26T17:57:58.265581Z","shell.execute_reply":"2023-06-26T17:57:58.265603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.fit(20, train_dataloader, val_dataloader, optimizer, criterion, scheduler=scheduler, early_stopper=early_stopper, start_epoch=29)","metadata":{"execution":{"iopub.status.busy":"2023-06-26T17:57:58.267471Z","iopub.status.idle":"2023-06-26T17:57:58.268353Z","shell.execute_reply.started":"2023-06-26T17:57:58.268080Z","shell.execute_reply":"2023-06-26T17:57:58.268104Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 5.Create submission","metadata":{}},{"cell_type":"code","source":"test_dataset = Dataset(test_files, test_transform, False)\ntest_dataloader = DataLoader(test_dataset, batch_size=128, shuffle=True)\nprint(f'Test dataset has {len(test_dataloader.dataset)} images')","metadata":{"execution":{"iopub.status.busy":"2023-06-26T17:57:58.269806Z","iopub.status.idle":"2023-06-26T17:57:58.270565Z","shell.execute_reply.started":"2023-06-26T17:57:58.270326Z","shell.execute_reply":"2023-06-26T17:57:58.270348Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.DataFrame(columns=['id', 'label'])\nfor image, id in test_dataloader:\n    output = model(image.to(device))\n    pred = torch.argmax(torch.nn.functional.softmax(output, dim=1), dim=1)\n    df = pd.concat([df, pd.DataFrame(zip(id, [t.item() for t in pred.cpu()]), columns=df.columns)], ignore_index=True)\ndf","metadata":{"execution":{"iopub.status.busy":"2023-06-26T17:57:58.272026Z","iopub.status.idle":"2023-06-26T17:57:58.272841Z","shell.execute_reply.started":"2023-06-26T17:57:58.272558Z","shell.execute_reply":"2023-06-26T17:57:58.272582Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.to_csv('submission.csv', index=False)","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-06-26T17:57:58.274237Z","iopub.status.idle":"2023-06-26T17:57:58.275043Z","shell.execute_reply.started":"2023-06-26T17:57:58.274790Z","shell.execute_reply":"2023-06-26T17:57:58.274813Z"},"trusted":true},"execution_count":null,"outputs":[]}]}