{"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":"# Without any aug","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:32:19.886502Z","iopub.execute_input":"2022-10-24T15:32:19.886908Z","iopub.status.idle":"2022-10-24T15:32:19.913700Z","shell.execute_reply.started":"2022-10-24T15:32:19.886807Z","shell.execute_reply":"2022-10-24T15:32:19.912905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install timm torch==1.7.0 torchvision==0.8.1\n!curl https://raw.githubusercontent.com/pytorch/xla/master/contrib/scripts/env-setup.py -o pytorch-xla-env-setup.py\n!python pytorch-xla-env-setup.py\n","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-10-24T15:32:19.915154Z","iopub.execute_input":"2022-10-24T15:32:19.916013Z","iopub.status.idle":"2022-10-24T15:34:18.415970Z","shell.execute_reply.started":"2022-10-24T15:32:19.915965Z","shell.execute_reply":"2022-10-24T15:34:18.414979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nplt.style.use(\"ggplot\")\n\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms as transforms\n\nimport torch_xla\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.xla_multiprocessing as xmp\nimport torch_xla.distributed.parallel_loader as pl\n\nimport timm\n\nimport gc\nimport os\nimport time\nimport random\nfrom datetime import datetime\n\nfrom PIL import Image\nfrom tqdm.notebook import tqdm\nfrom sklearn import model_selection, metrics\nfrom sklearn.model_selection import train_test_split","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:34:18.417712Z","iopub.execute_input":"2022-10-24T15:34:18.417990Z","iopub.status.idle":"2022-10-24T15:34:21.442406Z","shell.execute_reply.started":"2022-10-24T15:34:18.417956Z","shell.execute_reply":"2022-10-24T15:34:21.441510Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# For parallelization in TPUs\nos.environ[\"XLA_USE_BF16\"] = \"1\"\nos.environ[\"XLA_TENSOR_ALLOCATOR_MAXSIZE\"] = \"100000000\"","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:34:21.444046Z","iopub.execute_input":"2022-10-24T15:34:21.444314Z","iopub.status.idle":"2022-10-24T15:34:21.450176Z","shell.execute_reply.started":"2022-10-24T15:34:21.444283Z","shell.execute_reply":"2022-10-24T15:34:21.448654Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed):\n    \"\"\"\n    Seeds basic parameters for reproductibility of results\n    \n    Arguments:\n        seed {int} -- Number of the seed\n    \"\"\"\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\n\nseed_everything(1001)","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:34:21.452346Z","iopub.execute_input":"2022-10-24T15:34:21.452795Z","iopub.status.idle":"2022-10-24T15:34:21.466857Z","shell.execute_reply.started":"2022-10-24T15:34:21.452763Z","shell.execute_reply":"2022-10-24T15:34:21.465949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#make train, test sets\nfrom sklearn import preprocessing\n\nDATA_PATH = \"../input/sdss-images/data.csv\"\ndf = pd.read_csv(DATA_PATH)\n\nle = preprocessing.LabelEncoder()\ndf['class'] = le.fit_transform(df['class'])\n\ntrain,val = train_test_split(df,test_size=0.2,random_state=2022)\ntrain = train.reset_index().drop(['index'],axis=1)\nval = val.reset_index().drop(['index'],axis=1)\nval,test= train_test_split(val,test_size=0.5,random_state=2022)\nval = val.reset_index().drop(['index'],axis=1)\ntest = test.reset_index().drop(['index'],axis=1)","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:34:21.468387Z","iopub.execute_input":"2022-10-24T15:34:21.468672Z","iopub.status.idle":"2022-10-24T15:34:21.527532Z","shell.execute_reply.started":"2022-10-24T15:34:21.468641Z","shell.execute_reply":"2022-10-24T15:34:21.526863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"le.classes_","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:34:21.528590Z","iopub.execute_input":"2022-10-24T15:34:21.529169Z","iopub.status.idle":"2022-10-24T15:34:21.537992Z","shell.execute_reply.started":"2022-10-24T15:34:21.529135Z","shell.execute_reply":"2022-10-24T15:34:21.536990Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.head(2)","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:34:21.539244Z","iopub.execute_input":"2022-10-24T15:34:21.539484Z","iopub.status.idle":"2022-10-24T15:34:21.562710Z","shell.execute_reply.started":"2022-10-24T15:34:21.539456Z","shell.execute_reply":"2022-10-24T15:34:21.561950Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir images\n!mkdir images/val images/train","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:34:21.564119Z","iopub.execute_input":"2022-10-24T15:34:21.564986Z","iopub.status.idle":"2022-10-24T15:34:23.754740Z","shell.execute_reply.started":"2022-10-24T15:34:21.564926Z","shell.execute_reply":"2022-10-24T15:34:23.753433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import shutil\n\nfor img in train['image']:\n    shutil.copy(\"../input/sdss-images/images (1)/images/\"+img,\"./images/train/\"+img)\n    \nfor img in val['image']:\n    shutil.copy(\"../input/sdss-images/images (1)/images/\"+img,\"./images/val/\"+img)\n    ","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:34:23.757190Z","iopub.execute_input":"2022-10-24T15:34:23.757597Z","iopub.status.idle":"2022-10-24T15:35:01.909323Z","shell.execute_reply.started":"2022-10-24T15:34:23.757531Z","shell.execute_reply":"2022-10-24T15:35:01.908138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# general global variables\nDATA_PATH = \"../input/sdss-images/data.csv\"\nTRAIN_PATH = \"./images/train\"\nTEST_PATH = \"./images/val\"\nMODEL_PATH = (\n    \"../input/vit-base-models-pretrained-pytorch/jx_vit_base_p16_224-80ecf9dd.pth\"\n)\n\n# model specific global variables\nIMG_SIZE = 224\nBATCH_SIZE = 32 #32\nLR = 2e-05\nGAMMA = 0.7\nN_EPOCHS = 10\nwandb_args = {\"learning_rate\": LR, \"epochs\": N_EPOCHS, \"batch_size\": BATCH_SIZE,\"gamma\":GAMMA,\"img_size\":IMG_SIZE,\"project\":\"ViT SDSS\",\"name\":\"Default aug\"}","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:35:01.912162Z","iopub.execute_input":"2022-10-24T15:35:01.912425Z","iopub.status.idle":"2022-10-24T15:35:01.921006Z","shell.execute_reply.started":"2022-10-24T15:35:01.912394Z","shell.execute_reply":"2022-10-24T15:35:01.920015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SDSSDataset(torch.utils.data.Dataset):\n    \"\"\"\n    Helper Class to create the pytorch dataset\n    \"\"\"\n\n    def __init__(self, df, data_path=DATA_PATH, mode=\"train\", transforms=None):\n        super().__init__()\n        self.df_data = df.values\n        self.data_path = data_path\n        self.transforms = transforms\n        self.mode = mode\n        self.data_dir = \"./images/train\" if mode == \"train\" else \"./images/val\"\n\n    def __len__(self):\n        return len(self.df_data)\n\n    def __getitem__(self, index):\n        label,img_name = self.df_data[index]\n        img_path = os.path.join(self.data_dir, img_name)\n        img = Image.open(img_path).convert(\"RGB\")\n\n        if self.transforms is not None:\n            image = self.transforms(img)\n\n        return image, label","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:35:01.922388Z","iopub.execute_input":"2022-10-24T15:35:01.922705Z","iopub.status.idle":"2022-10-24T15:35:01.932431Z","shell.execute_reply.started":"2022-10-24T15:35:01.922673Z","shell.execute_reply":"2022-10-24T15:35:01.931484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# create image augmentations\ntransforms_train = transforms.Compose(\n    [\n        transforms.Resize((IMG_SIZE, IMG_SIZE)),\n#         transforms.RandomHorizontalFlip(p=0.3),\n#         transforms.RandomVerticalFlip(p=0.3),\n        transforms.RandomResizedCrop(IMG_SIZE),\n        transforms.ToTensor(),\n#         transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),\n    ]\n)\n\ntransforms_valid = transforms.Compose(\n    [\n        transforms.Resize((IMG_SIZE, IMG_SIZE)),\n        transforms.ToTensor(),\n#         transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),\n    ]\n)","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:35:01.933657Z","iopub.execute_input":"2022-10-24T15:35:01.934000Z","iopub.status.idle":"2022-10-24T15:35:01.946196Z","shell.execute_reply.started":"2022-10-24T15:35:01.933972Z","shell.execute_reply":"2022-10-24T15:35:01.945246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Available Vision Transformer Models: \")\ntimm.list_models(\"vit*\")","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:35:01.949606Z","iopub.execute_input":"2022-10-24T15:35:01.950610Z","iopub.status.idle":"2022-10-24T15:35:01.970351Z","shell.execute_reply.started":"2022-10-24T15:35:01.950564Z","shell.execute_reply":"2022-10-24T15:35:01.969312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ViTBase16(nn.Module):\n    def __init__(self, n_classes, pretrained=False):\n\n        super(ViTBase16, self).__init__()\n\n        self.model = timm.create_model(\"vit_base_patch16_224\", pretrained=False)\n        if pretrained:\n            self.model.load_state_dict(torch.load(MODEL_PATH))\n\n        self.model.head = nn.Linear(self.model.head.in_features, n_classes)\n\n    def forward(self, x):\n        x = self.model(x)\n        return x\n\n    def train_one_epoch(self, train_loader, criterion, optimizer, device):\n        # keep track of training loss\n        epoch_loss = 0.0\n        epoch_accuracy = 0.0\n\n        ###################\n        # train the model #\n        ###################\n        self.model.train()\n        for i, (data, target) in enumerate(train_loader):\n            # move tensors to GPU if CUDA is available\n            if device.type == \"cuda\":\n                data, target = data.cuda(), target.cuda()\n            elif device.type == \"xla\":\n                data = data.to(device, dtype=torch.float32)\n#                 print(target)\n                target = torch.tensor(target).to(device, dtype=torch.int64)\n\n            # clear the gradients of all optimized variables\n            optimizer.zero_grad()\n            # forward pass: compute predicted outputs by passing inputs to the model\n            output = self.forward(data)\n            # calculate the batch loss\n            loss = criterion(output, target)\n            # backward pass: compute gradient of the loss with respect to model parameters\n            loss.backward()\n            # Calculate Accuracy\n            accuracy = (output.argmax(dim=1) == target).float().mean()\n            # update training loss and accuracy\n            epoch_loss += loss\n            epoch_accuracy += accuracy\n\n            # perform a single optimization step (parameter update)\n            if device.type == \"xla\":\n                xm.optimizer_step(optimizer)\n\n                if i % 20 == 0:\n                    xm.master_print(f\"\\tBATCH {i+1}/{len(train_loader)} - LOSS: {loss}\")\n\n            else:\n                optimizer.step()\n\n        return epoch_loss / len(train_loader), epoch_accuracy / len(train_loader)\n\n    def validate_one_epoch(self, valid_loader, criterion, device):\n        # keep track of validation loss\n        valid_loss = 0.0\n        valid_accuracy = 0.0\n\n        ######################\n        # validate the model #\n        ######################\n        self.model.eval()\n        for data, target in valid_loader:\n            # move tensors to GPU if CUDA is available\n            if device.type == \"cuda\":\n                data, target = data.cuda(), target.cuda()\n            elif device.type == \"xla\":\n                data = data.to(device, dtype=torch.float32)\n                target = torch.tensor(target).to(device, dtype=torch.int64)\n\n            with torch.no_grad():\n                # forward pass: compute predicted outputs by passing inputs to the model\n                output = self.model(data)\n                # calculate the batch loss\n                loss = criterion(output, target)\n                # Calculate Accuracy\n                accuracy = (output.argmax(dim=1) == target).float().mean()\n                # update average validation loss and accuracy\n                valid_loss += loss\n                valid_accuracy += accuracy\n\n        return valid_loss / len(valid_loader), valid_accuracy / len(valid_loader)","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:35:01.971827Z","iopub.execute_input":"2022-10-24T15:35:01.972092Z","iopub.status.idle":"2022-10-24T15:35:01.990917Z","shell.execute_reply.started":"2022-10-24T15:35:01.972064Z","shell.execute_reply":"2022-10-24T15:35:01.989859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import wandb\nimport os\nfrom kaggle_secrets import UserSecretsClient\n\n# user_secrets = UserSecretsClient()\n# secret_value_0 = user_secrets.get_secret(\"wandb_key\")\n\n# os.environ[\"WANDB_API_KEY\"] = secret_value_0\n# os.environ[\"WANDB_MODE\"] = \"offline\"\n\nmodel = ViTBase16(n_classes=3, pretrained=True)\n# wandb.init(config=wandb_args)\n\n# wandb.watch(model, log_freq=100)\n\ndef fit_tpu(\n    model, epochs, device, criterion, optimizer, train_loader, valid_loader=None\n):\n\n    valid_loss_min = np.Inf  # track change in validation loss\n\n    # keeping track of losses as it happen\n    train_losses = []\n    valid_losses = []\n    train_accs = []\n    valid_accs = []\n    \n    \n\n\n    for epoch in range(1, epochs + 1):\n        gc.collect()\n        para_train_loader = pl.ParallelLoader(train_loader, [device])\n        \n        xm.master_print(f\"{'='*50}\")\n        xm.master_print(f\"EPOCH {epoch} - TRAINING...\")\n        train_loss, train_acc = model.train_one_epoch(\n            para_train_loader.per_device_loader(device), criterion, optimizer, device\n        )\n        xm.master_print(\n            f\"\\n\\t[TRAIN] EPOCH {epoch} - LOSS: {train_loss}, ACCURACY: {train_acc}\\n\"\n        )\n        train_losses.append(train_loss)\n        train_accs.append(train_acc)\n        gc.collect()\n        \n        if valid_loader is not None:\n            gc.collect()\n            para_valid_loader = pl.ParallelLoader(valid_loader, [device])\n            xm.master_print(f\"EPOCH {epoch} - VALIDATING...\")\n            valid_loss, valid_acc = model.validate_one_epoch(\n                para_valid_loader.per_device_loader(device), criterion, device\n            )\n            xm.master_print(f\"\\t[VALID] LOSS: {valid_loss}, ACCURACY: {valid_acc}\\n\")\n            valid_losses.append(valid_loss)\n            valid_accs.append(valid_acc)\n            gc.collect()\n\n            # save model if validation loss has decreased\n            if valid_loss <= valid_loss_min and epoch != 1:\n                xm.master_print(\n                    \"Validation loss decreased ({:.4f} --> {:.4f}).  Saving model ...\".format(\n                        valid_loss_min, valid_loss\n                    )\n                )\n            #                 xm.save(model.state_dict(), 'best_model.pth')\n\n            valid_loss_min = valid_loss\n    logs={\n        \"train_loss\": train_losses,\n        \"valid_losses\": valid_losses,\n        \"train_acc\": train_accs,\n        \"valid_acc\": valid_accs,\n    }\n#     wandb.log(logs)\n    return logs","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:35:01.992476Z","iopub.execute_input":"2022-10-24T15:35:01.992785Z","iopub.status.idle":"2022-10-24T15:35:06.757793Z","shell.execute_reply.started":"2022-10-24T15:35:01.992752Z","shell.execute_reply":"2022-10-24T15:35:06.756959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _run():\n    train_dataset = SDSSDataset(train, transforms=transforms_train)\n    valid_dataset = SDSSDataset(val, transforms=transforms_valid,mode=\"valid\",)\n\n    train_sampler = torch.utils.data.distributed.DistributedSampler(\n        train_dataset,\n        num_replicas=xm.xrt_world_size(),\n        rank=xm.get_ordinal(),\n        shuffle=True,\n    )\n\n    valid_sampler = torch.utils.data.distributed.DistributedSampler(\n        valid_dataset,\n        num_replicas=xm.xrt_world_size(),\n        rank=xm.get_ordinal(),\n        shuffle=False,\n    )\n\n    train_loader = torch.utils.data.DataLoader(\n        dataset=train_dataset,\n        batch_size=BATCH_SIZE,\n        sampler=train_sampler,\n        drop_last=True,\n        num_workers=8,\n    )\n\n    valid_loader = torch.utils.data.DataLoader(\n        dataset=valid_dataset,\n        batch_size=BATCH_SIZE,\n        sampler=valid_sampler,\n        drop_last=True,\n        num_workers=8,\n    )\n\n    criterion = nn.CrossEntropyLoss()\n    #     device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    device = xm.xla_device()\n    model.to(device)\n    \n    lr = LR * xm.xrt_world_size()\n    optimizer = torch.optim.Adam(model.parameters(), lr=lr)\n\n    xm.master_print(f\"INITIALIZING TRAINING ON {xm.xrt_world_size()} TPU CORES\")\n    start_time = datetime.now()\n    xm.master_print(f\"Start Time: {start_time}\")\n\n    logs = fit_tpu(\n        model=model,\n        epochs=N_EPOCHS,\n        device=device,\n        criterion=criterion,\n        optimizer=optimizer,\n        train_loader=train_loader,\n        valid_loader=valid_loader,\n    )\n#     wandb.log(logs)\n    xm.master_print(f\"Execution time: {datetime.now() - start_time}\")\n\n    xm.master_print(\"Saving Model\")\n    xm.save(\n        model.state_dict(), f'model_5e_{datetime.now().strftime(\"%Y%m%d-%H%M\")}.pth'\n    )","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:35:06.758865Z","iopub.execute_input":"2022-10-24T15:35:06.759329Z","iopub.status.idle":"2022-10-24T15:35:06.770943Z","shell.execute_reply.started":"2022-10-24T15:35:06.759271Z","shell.execute_reply":"2022-10-24T15:35:06.769986Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n# Start training processes\ndef _mp_fn(rank, flags):\n    torch.set_default_tensor_type(\"torch.FloatTensor\")\n    a = _run()\n\n\n# _run()\nFLAGS = {}\nxmp.spawn(_mp_fn, args=(FLAGS,), nprocs=8, start_method=\"fork\")","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:35:06.772776Z","iopub.execute_input":"2022-10-24T15:35:06.773374Z","iopub.status.idle":"2022-10-24T15:42:23.889513Z","shell.execute_reply.started":"2022-10-24T15:35:06.773309Z","shell.execute_reply":"2022-10-24T15:42:23.888101Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:42:23.892191Z","iopub.execute_input":"2022-10-24T15:42:23.892580Z","iopub.status.idle":"2022-10-24T15:42:25.032506Z","shell.execute_reply.started":"2022-10-24T15:42:23.892509Z","shell.execute_reply":"2022-10-24T15:42:25.031430Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -r images","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:42:25.034690Z","iopub.execute_input":"2022-10-24T15:42:25.035085Z","iopub.status.idle":"2022-10-24T15:42:26.414229Z","shell.execute_reply.started":"2022-10-24T15:42:25.035038Z","shell.execute_reply":"2022-10-24T15:42:26.412659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"!ls","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:42:26.416659Z","iopub.execute_input":"2022-10-24T15:42:26.416970Z","iopub.status.idle":"2022-10-24T15:42:27.538754Z","shell.execute_reply.started":"2022-10-24T15:42:26.416936Z","shell.execute_reply":"2022-10-24T15:42:27.537714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_inference = ViTBase16(n_classes=3, pretrained=True)","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:42:27.541027Z","iopub.execute_input":"2022-10-24T15:42:27.542025Z","iopub.status.idle":"2022-10-24T15:42:32.375967Z","shell.execute_reply.started":"2022-10-24T15:42:27.541975Z","shell.execute_reply":"2022-10-24T15:42:32.375106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_inference.load_state_dict(torch.load(\"./model_5e_20221024-1542.pth\"))","metadata":{"execution":{"iopub.status.busy":"2022-10-24T15:57:36.905279Z","iopub.execute_input":"2022-10-24T15:57:36.906054Z","iopub.status.idle":"2022-10-24T15:57:37.814227Z","shell.execute_reply.started":"2022-10-24T15:57:36.906013Z","shell.execute_reply":"2022-10-24T15:57:37.813166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SDSSDataset(torch.utils.data.Dataset):\n    \"\"\"\n    Helper Class to create the pytorch dataset\n    \"\"\"\n\n    def __init__(self, df, data_path=DATA_PATH, transforms=None):\n        super().__init__()\n        self.df_data = df.values\n        self.data_path = data_path\n        self.transforms = transforms\n        \n        self.data_dir = \"./images/test\"\n\n    def __len__(self):\n        return len(self.df_data)\n\n    def __getitem__(self, index):\n        label,img_name = self.df_data[index]\n        img_path = os.path.join(self.data_dir, img_name)\n        img = Image.open(img_path).convert(\"RGB\")\n\n        if self.transforms is not None:\n            image = self.transforms(img)\n\n        return image, label","metadata":{"execution":{"iopub.status.busy":"2022-10-24T16:11:27.108450Z","iopub.execute_input":"2022-10-24T16:11:27.109385Z","iopub.status.idle":"2022-10-24T16:11:27.121833Z","shell.execute_reply.started":"2022-10-24T16:11:27.109314Z","shell.execute_reply":"2022-10-24T16:11:27.120076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms_valid = transforms.Compose(\n    [\n        transforms.Resize((IMG_SIZE, IMG_SIZE)),\n        transforms.ToTensor(),\n    ]\n)\n\ntest_dataset = SDSSDataset(test, transforms=transforms_valid)\ntest_sampler = torch.utils.data.distributed.DistributedSampler(\n        test_dataset,\n        num_replicas=xm.xrt_world_size(),\n        rank=xm.get_ordinal(),\n        shuffle=False,\n    )\n\ntest_loader = torch.utils.data.DataLoader(\ndataset=test_dataset,\nbatch_size=1,\nsampler=test_sampler,\ndrop_last=True,\nnum_workers=8,\n\n)","metadata":{"execution":{"iopub.status.busy":"2022-10-24T16:11:27.362930Z","iopub.execute_input":"2022-10-24T16:11:27.363239Z","iopub.status.idle":"2022-10-24T16:11:27.372415Z","shell.execute_reply.started":"2022-10-24T16:11:27.363209Z","shell.execute_reply":"2022-10-24T16:11:27.371332Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -r images\n!mkdir images\n!mkdir images/test","metadata":{"execution":{"iopub.status.busy":"2022-10-24T16:11:55.958923Z","iopub.execute_input":"2022-10-24T16:11:55.959228Z","iopub.status.idle":"2022-10-24T16:11:59.462142Z","shell.execute_reply.started":"2022-10-24T16:11:55.959200Z","shell.execute_reply":"2022-10-24T16:11:59.460713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import shutil\n\nfor img in test['image']:\n    shutil.copy(\"../input/sdss-images/images (1)/images/\"+img,\"./images/test/\"+img)\n    ","metadata":{"execution":{"iopub.status.busy":"2022-10-24T16:11:59.466174Z","iopub.execute_input":"2022-10-24T16:11:59.466622Z","iopub.status.idle":"2022-10-24T16:12:05.488010Z","shell.execute_reply.started":"2022-10-24T16:11:59.466569Z","shell.execute_reply":"2022-10-24T16:12:05.486895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for x in test_loader:\n    print(len(x))\n    break","metadata":{"execution":{"iopub.status.busy":"2022-10-24T16:12:11.652701Z","iopub.execute_input":"2022-10-24T16:12:11.653043Z","iopub.status.idle":"2022-10-24T16:12:12.125182Z","shell.execute_reply.started":"2022-10-24T16:12:11.653008Z","shell.execute_reply":"2022-10-24T16:12:12.122642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_inf = model_inference.eval()","metadata":{"execution":{"iopub.status.busy":"2022-10-24T16:12:14.633439Z","iopub.execute_input":"2022-10-24T16:12:14.633831Z","iopub.status.idle":"2022-10-24T16:12:14.644600Z","shell.execute_reply.started":"2022-10-24T16:12:14.633794Z","shell.execute_reply":"2022-10-24T16:12:14.643161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_inf(x[0])","metadata":{"execution":{"iopub.status.busy":"2022-10-24T16:12:15.540420Z","iopub.execute_input":"2022-10-24T16:12:15.540939Z","iopub.status.idle":"2022-10-24T16:12:16.701271Z","shell.execute_reply.started":"2022-10-24T16:12:15.540883Z","shell.execute_reply":"2022-10-24T16:12:16.700528Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = xm.xla_device()\n\nmodel_inf.to(device)\nprint(\"moved\")","metadata":{"execution":{"iopub.status.busy":"2022-10-24T16:12:18.656978Z","iopub.execute_input":"2022-10-24T16:12:18.657661Z","iopub.status.idle":"2022-10-24T16:12:25.426519Z","shell.execute_reply.started":"2022-10-24T16:12:18.657592Z","shell.execute_reply":"2022-10-24T16:12:25.425397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_preds = list()\n\n# device = xm.xla_device()\n# model.to(device)\n\nfor x in tqdm(test_loader):\n    test_preds.append(model_inf(x[0].to(device)))","metadata":{"execution":{"iopub.status.busy":"2022-10-24T16:12:31.494511Z","iopub.execute_input":"2022-10-24T16:12:31.494874Z","iopub.status.idle":"2022-10-24T16:12:56.041843Z","shell.execute_reply.started":"2022-10-24T16:12:31.494839Z","shell.execute_reply":"2022-10-24T16:12:56.040151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_preds = [list(x.cpu().detach().numpy()[0]) for x in test_preds]","metadata":{"execution":{"iopub.status.busy":"2022-10-24T16:13:05.486112Z","iopub.execute_input":"2022-10-24T16:13:05.486752Z","iopub.status.idle":"2022-10-24T16:13:20.940467Z","shell.execute_reply.started":"2022-10-24T16:13:05.486698Z","shell.execute_reply":"2022-10-24T16:13:20.939385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_labels = [np.argmax(x) for x in test_preds]","metadata":{"execution":{"iopub.status.busy":"2022-10-24T16:13:21.028902Z","iopub.execute_input":"2022-10-24T16:13:21.030229Z","iopub.status.idle":"2022-10-24T16:13:21.046147Z","shell.execute_reply.started":"2022-10-24T16:13:21.030169Z","shell.execute_reply":"2022-10-24T16:13:21.045131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# val_labels","metadata":{"execution":{"iopub.status.busy":"2022-10-24T16:13:21.047836Z","iopub.execute_input":"2022-10-24T16:13:21.048373Z","iopub.status.idle":"2022-10-24T16:13:21.054753Z","shell.execute_reply.started":"2022-10-24T16:13:21.048333Z","shell.execute_reply":"2022-10-24T16:13:21.053498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"true_labels = list()\nfor x in test_loader:\n    true_labels.append(x[1].detach().cpu().numpy()[0])","metadata":{"execution":{"iopub.status.busy":"2022-10-24T16:13:21.443686Z","iopub.execute_input":"2022-10-24T16:13:21.444015Z","iopub.status.idle":"2022-10-24T16:13:25.122162Z","shell.execute_reply.started":"2022-10-24T16:13:21.443983Z","shell.execute_reply":"2022-10-24T16:13:25.120901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import classification_report","metadata":{"execution":{"iopub.status.busy":"2022-10-24T16:13:27.453401Z","iopub.execute_input":"2022-10-24T16:13:27.454644Z","iopub.status.idle":"2022-10-24T16:13:27.459405Z","shell.execute_reply.started":"2022-10-24T16:13:27.454592Z","shell.execute_reply":"2022-10-24T16:13:27.458402Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_labels[:10],true_labels[:10]","metadata":{"execution":{"iopub.status.busy":"2022-10-24T16:13:32.165377Z","iopub.execute_input":"2022-10-24T16:13:32.166705Z","iopub.status.idle":"2022-10-24T16:13:32.174543Z","shell.execute_reply.started":"2022-10-24T16:13:32.166651Z","shell.execute_reply":"2022-10-24T16:13:32.173453Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"classification_report(test_labels,true_labels).split(\"\\n\")","metadata":{"execution":{"iopub.status.busy":"2022-10-24T16:13:38.481803Z","iopub.execute_input":"2022-10-24T16:13:38.482176Z","iopub.status.idle":"2022-10-24T16:13:38.501145Z","shell.execute_reply.started":"2022-10-24T16:13:38.482137Z","shell.execute_reply":"2022-10-24T16:13:38.500087Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}