{"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":"Hi, I'm a beginner and now learning Variational Autoencoder with Pytorch, but it seems that it's not working well.\n\nHere's a link to the discussion page!\n\nhttps://www.kaggle.com/c/h-and-m-personalized-fashion-recommendations/discussion/307169\n\nWhat improvements are needed for successful learning?\n\nfor example,\n\nchange loss function speed up and train with more epochs (=> How can we speed up the process?) If the implementation of vae is successful, I think we can incorporate innovations such as recommending closer vectors in the latent space!\n\nAny advice would be appreciated!","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n%matplotlib inline\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\n\nimport torchvision\nfrom torchvision import transforms\nfrom torchvision.utils import make_grid\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.datasets import FashionMNIST\n\nimport os\nimport cv2\nfrom PIL import Image\nfrom tqdm import tqdm\n\n# import torchbearer\n# import torchbearer.callbacks as callbacks\n# from torchbearer import Trial, state_key\n\n# MU = state_key('mu')\n# LOGVAR = state_key('logvar')","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:07:57.285251Z","iopub.execute_input":"2022-02-13T02:07:57.285509Z","iopub.status.idle":"2022-02-13T02:07:57.295158Z","shell.execute_reply.started":"2022-02-13T02:07:57.285479Z","shell.execute_reply":"2022-02-13T02:07:57.294435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trs_data = pd.read_csv('../input/h-and-m-personalized-fashion-recommendations/transactions_train.csv')\ntrs_data","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"article_data = pd.read_csv('../input/h-and-m-personalized-fashion-recommendations/articles.csv')\nlen(article_data)","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:08:04.421127Z","iopub.execute_input":"2022-02-13T02:08:04.421656Z","iopub.status.idle":"2022-02-13T02:08:05.462190Z","shell.execute_reply.started":"2022-02-13T02:08:04.421618Z","shell.execute_reply":"2022-02-13T02:08:05.461411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## 画像がない品物の情報を、article_dataから取り除く\nfor i in tqdm(range(105542)):\n    base_path = '../input/h-and-m-personalized-fashion-recommendations/images'\n    article_id = '0' + str(article_data.loc[i]['article_id']) + '.jpg'\n    image_path = os.path.join(base_path, article_id[0:3], article_id)\n    if (not os.path.exists(image_path)):\n        article_data = article_data.drop(i)\n        \n    \nlen(article_data) #105100","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:08:05.463687Z","iopub.execute_input":"2022-02-13T02:08:05.464007Z","iopub.status.idle":"2022-02-13T02:10:28.072523Z","shell.execute_reply.started":"2022-02-13T02:08:05.463969Z","shell.execute_reply":"2022-02-13T02:10:28.071833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"article_data = article_data.reset_index(drop=True)\narticle_data","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:10:28.074171Z","iopub.execute_input":"2022-02-13T02:10:28.074595Z","iopub.status.idle":"2022-02-13T02:10:28.181451Z","shell.execute_reply.started":"2022-02-13T02:10:28.074545Z","shell.execute_reply":"2022-02-13T02:10:28.180630Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ## 25782行目の画像は、(1,28,28)になっている\n# for i in range(len(image_dataset)):\n#     if (image_dataset[i].shape[0]!=3):\n#         print(i)\n    \n#     if i%2000==0:\n#         print(\"now{}\".format(i))","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:10:28.182896Z","iopub.execute_input":"2022-02-13T02:10:28.183342Z","iopub.status.idle":"2022-02-13T02:10:28.187856Z","shell.execute_reply.started":"2022-02-13T02:10:28.183300Z","shell.execute_reply":"2022-02-13T02:10:28.186901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"article_data = article_data.drop(25782)\narticle_data = article_data.reset_index(drop=True)\narticle_data","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:10:28.190066Z","iopub.execute_input":"2022-02-13T02:10:28.190426Z","iopub.status.idle":"2022-02-13T02:10:28.331941Z","shell.execute_reply.started":"2022-02-13T02:10:28.190384Z","shell.execute_reply":"2022-02-13T02:10:28.331038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_img_path(id):\n    article_id = '0' + str(id) + '.jpg'\n    base_path = \"../input/h-and-m-personalized-fashion-recommendations/images\"\n    image_path = os.path.join(base_path, article_id[0:3], article_id)\n    return image_path\n\narticle_data['path'] = article_data['article_id'].map(get_img_path)\nimg_path_list = article_data['path']\nimg_path_list[0]","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:10:28.333516Z","iopub.execute_input":"2022-02-13T02:10:28.334110Z","iopub.status.idle":"2022-02-13T02:10:28.642964Z","shell.execute_reply.started":"2022-02-13T02:10:28.334066Z","shell.execute_reply":"2022-02-13T02:10:28.642154Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import glob\n\n# image_path_list = glob.glob(\"../input/h-and-m-personalized-fashion-recommendations/images/**/*.jpg\", recursive=True)\n# len(image_path_list)\n\n# image_path_list[0]","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:10:28.644419Z","iopub.execute_input":"2022-02-13T02:10:28.644693Z","iopub.status.idle":"2022-02-13T02:10:28.648192Z","shell.execute_reply.started":"2022-02-13T02:10:28.644658Z","shell.execute_reply":"2022-02-13T02:10:28.647359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class fashion_Datasets(Dataset):\n    def __init__(self, img_path_list, transform=None):\n        self.img_path_list = img_path_list\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.img_path_list)\n    \n    def __getitem__(self, i):\n        img = Image.open(self.img_path_list[i])\n        \n        if self.transform is not None:\n            img = self.transform(img)\n        \n        return img","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:10:28.649592Z","iopub.execute_input":"2022-02-13T02:10:28.649834Z","iopub.status.idle":"2022-02-13T02:10:28.658245Z","shell.execute_reply.started":"2022-02-13T02:10:28.649802Z","shell.execute_reply":"2022-02-13T02:10:28.657603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform =  transforms.Compose([transforms.Resize((28,28)), transforms.ToTensor()])\nimage_dataset = fashion_Datasets(img_path_list = img_path_list, transform=transform)","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:10:28.659033Z","iopub.execute_input":"2022-02-13T02:10:28.660349Z","iopub.status.idle":"2022-02-13T02:10:28.666674Z","shell.execute_reply.started":"2022-02-13T02:10:28.660312Z","shell.execute_reply":"2022-02-13T02:10:28.665884Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for i in range(len(image_dataset)):\n#     if (image_dataset[i].shape[0]!=3):\n#         print(i)\n    \n#     if i%2000==0:\n#         print(\"now{}\".format(i))","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:10:28.668003Z","iopub.execute_input":"2022-02-13T02:10:28.668303Z","iopub.status.idle":"2022-02-13T02:10:28.675337Z","shell.execute_reply.started":"2022-02-13T02:10:28.668263Z","shell.execute_reply":"2022-02-13T02:10:28.674651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# \"0616100001.jpg\"","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:10:28.678047Z","iopub.execute_input":"2022-02-13T02:10:28.678325Z","iopub.status.idle":"2022-02-13T02:10:28.684023Z","shell.execute_reply.started":"2022-02-13T02:10:28.678270Z","shell.execute_reply":"2022-02-13T02:10:28.683258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# article_data = article_data.drop(25782)\n# article_data = article_data.reset_index(drop=True)\n# article_data","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:10:28.685006Z","iopub.execute_input":"2022-02-13T02:10:28.685259Z","iopub.status.idle":"2022-02-13T02:10:28.692313Z","shell.execute_reply.started":"2022-02-13T02:10:28.685216Z","shell.execute_reply":"2022-02-13T02:10:28.691589Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_size = 5000\ntrain_size = 20000\nothers = len(image_dataset) - train_size - val_size\ntrain_data, val_data, _ = torch.utils.data.random_split(image_dataset, [train_size, val_size, others])","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:10:28.693138Z","iopub.execute_input":"2022-02-13T02:10:28.693325Z","iopub.status.idle":"2022-02-13T02:10:28.719438Z","shell.execute_reply.started":"2022-02-13T02:10:28.693294Z","shell.execute_reply":"2022-02-13T02:10:28.718850Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for i in range(10000):\n#     k = train_data[i]","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:10:28.720601Z","iopub.execute_input":"2022-02-13T02:10:28.720844Z","iopub.status.idle":"2022-02-13T02:10:28.723954Z","shell.execute_reply.started":"2022-02-13T02:10:28.720812Z","shell.execute_reply":"2022-02-13T02:10:28.723339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(val_data)","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:10:28.725372Z","iopub.execute_input":"2022-02-13T02:10:28.725901Z","iopub.status.idle":"2022-02-13T02:10:28.734283Z","shell.execute_reply.started":"2022-02-13T02:10:28.725863Z","shell.execute_reply":"2022-02-13T02:10:28.733593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataloader_train = DataLoader(train_data, batch_size=64, shuffle=True, num_workers=2)\ndataloader_valid = DataLoader(val_data, batch_size=64, shuffle=True, num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:10:28.735608Z","iopub.execute_input":"2022-02-13T02:10:28.735854Z","iopub.status.idle":"2022-02-13T02:10:28.743461Z","shell.execute_reply.started":"2022-02-13T02:10:28.735821Z","shell.execute_reply":"2022-02-13T02:10:28.742724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# torch.log(0)によるnanを防ぐ\ndef torch_log(x):\n    return torch.log(torch.clamp(x, min=1e-10))","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:10:28.744732Z","iopub.execute_input":"2022-02-13T02:10:28.744972Z","iopub.status.idle":"2022-02-13T02:10:28.752318Z","shell.execute_reply.started":"2022-02-13T02:10:28.744939Z","shell.execute_reply":"2022-02-13T02:10:28.751611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class VAE(nn.Module):\n    def __init__(self, latent_size):\n        super(VAE, self).__init__()\n        self.latent_size = latent_size\n\n        self.encoder = nn.Sequential(\n            nn.Conv2d(3, 32, 3, 2, 1),  \n            nn.ReLU(True),\n            nn.Conv2d(32, 64, 3, 2, 1), \n            nn.ReLU(True),\n            nn.Conv2d(64, 64, 3, 2, 1),  \n            \n        )\n        \n        self.mu = nn.Linear(64 * 4 * 4, latent_size)\n        self.logvar = nn.Linear(64 * 4 * 4, latent_size)\n        \n        self.upsample = nn.Linear(latent_size, 64 * 3 * 3)\n        self.decoder = nn.Sequential(\n            nn.ConvTranspose2d(64, 64, kernel_size=3, stride=2),\n            nn.ReLU(True),\n            nn.ConvTranspose2d(64, 32, kernel_size=2, stride=2),\n            nn.ReLU(True),\n            nn.ConvTranspose2d(32, 3, kernel_size=2, stride=2)  \n        )\n\n    def reparameterize(self, mu, logvar):\n        if self.training:\n            std = torch.exp(0.5*logvar)\n            eps = torch.randn_like(std)\n#             print(\"training\")\n            return eps.mul(std).add_(mu)\n        else:\n#             print(\"not training\")\n            return mu\n\n    def forward(self, x):\n        image = x\n        x = self.encoder(x).relu().view(x.size(0), -1)\n        \n        mu = self.mu(x)\n        logvar = self.logvar(x)\n        z = self.reparameterize(mu, logvar)\n        \n        result = self.decoder(self.upsample(z).relu().view(-1, 64, 3, 3))\n        \n#         if state is not None:\n#             state[torchbearer.Y_TRUE] = image\n#             state[MU] = mu\n#             state[LOGVAR] = logvar\n        \n        return result, z\n    \n    def loss(self, x):\n        \n#         print('x.shape')\n#         print(x.shape)\n#         print(self.encoder(x).relu().shape)\n        out = self.encoder(x).relu().view(x.size(0), -1) #(64,3,28,28) -> (64,64,4,4) -> (64, 1024)\n#         print('out.shape')\n#         print(out.shape)\n#         print('loss')\n        mu = self.mu(out)\n#         print('mu.shape')\n#         print(mu.shape)\n        log_var = self.logvar(out)\n#         print('log_var.shape')\n#         print(log_var.shape)\n        \n        KL = -0.5 * torch.mean(torch.sum(1 + log_var - mu**2 - torch.exp(log_var), dim=1))\n        \n        z = self.reparameterize(mu, log_var)\n#         print(self.upsample(z).relu().shape)\n#         print(\"upsample\")\n#         print(z.shape)\n        y = self.decoder(self.upsample(z).relu().view(-1,64,3,3))\n        y = y.view(-1, 3*28*28)\n        x1 = x.view(-1, 3*28*28)\n        \n        reconstruction = torch.mean(torch.sum(x1 * torch_log(y) + (1-x1) * torch_log(1-y), dim=1))\n        \n        return KL, -reconstruction","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:10:28.755346Z","iopub.execute_input":"2022-02-13T02:10:28.755534Z","iopub.status.idle":"2022-02-13T02:10:28.770887Z","shell.execute_reply.started":"2022-02-13T02:10:28.755512Z","shell.execute_reply":"2022-02-13T02:10:28.770149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"z_dim = 20\nn_epochs = 10\nlr = 0.001\ndevice = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\nmodel = VAE(z_dim).to(device)","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:10:28.772377Z","iopub.execute_input":"2022-02-13T02:10:28.772882Z","iopub.status.idle":"2022-02-13T02:10:31.472956Z","shell.execute_reply.started":"2022-02-13T02:10:28.772845Z","shell.execute_reply":"2022-02-13T02:10:31.472214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(next(model.parameters()).is_cuda)","metadata":{"execution":{"iopub.status.busy":"2022-02-13T02:10:31.474587Z","iopub.execute_input":"2022-02-13T02:10:31.474875Z","iopub.status.idle":"2022-02-13T02:10:31.482267Z","shell.execute_reply.started":"2022-02-13T02:10:31.474842Z","shell.execute_reply":"2022-02-13T02:10:31.481546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = optim.Adam(model.parameters(), lr=lr)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n_epoch=5","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in range(n_epochs):\n    if (epoch%9==0):\n      lr /= 5\n    losses = []\n    KL_losses = []\n    reconstruction_losses = []\n    model.train()\n    for x in tqdm(dataloader_train):\n        x = x.to(device)\n\n        model.zero_grad()\n\n        y = model(x)\n\n        KL_loss, reconstruction_loss = model.loss(x)\n\n        loss = KL_loss + reconstruction_loss\n\n        loss.backward()\n        optimizer.step()\n\n        losses.append(loss.cpu().detach().numpy())\n        KL_losses.append(KL_loss.cpu().detach().numpy())\n        reconstruction_losses.append(reconstruction_loss.cpu().detach().numpy())\n\n    losses_val = []\n    model.eval()\n    for x in tqdm(dataloader_valid):\n\n        x = x.to(device)\n\n        y = model(x)\n\n        KL_loss, reconstruction_loss = model.loss(x)\n\n        loss = KL_loss + reconstruction_loss\n\n        # WRITE ME\n\n        losses_val.append(loss.cpu().detach().numpy())\n\n    print('EPOCH:%d, Train Lower Bound:%lf, (%lf, %lf), Valid Lower Bound:%lf' %\n          (epoch+1, np.average(losses), np.average(KL_losses), np.average(reconstruction_losses), np.average(losses_val)))","metadata":{"execution":{"iopub.status.busy":"2022-02-13T04:10:51.742674Z","iopub.execute_input":"2022-02-13T04:10:51.743159Z","iopub.status.idle":"2022-02-13T06:03:26.346769Z","shell.execute_reply.started":"2022-02-13T04:10:51.743123Z","shell.execute_reply":"2022-02-13T06:03:26.345943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.eval()\nfor x in tqdm(dataloader_valid):\n\n    x = x.to(device)\n\n    y, z = model(x)\n\n    break","metadata":{"execution":{"iopub.status.busy":"2022-02-13T06:04:38.248054Z","iopub.execute_input":"2022-02-13T06:04:38.248366Z","iopub.status.idle":"2022-02-13T06:04:45.213838Z","shell.execute_reply.started":"2022-02-13T06:04:38.248329Z","shell.execute_reply":"2022-02-13T06:04:45.212547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"original_im = train_data[0].numpy()\noriginal_im = np.transpose(original_im, (1,2,0))\nplt.imshow(original_im)","metadata":{"execution":{"iopub.status.busy":"2022-02-13T06:04:55.170621Z","iopub.execute_input":"2022-02-13T06:04:55.170911Z","iopub.status.idle":"2022-02-13T06:04:55.384095Z","shell.execute_reply.started":"2022-02-13T06:04:55.170878Z","shell.execute_reply":"2022-02-13T06:04:55.383401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x = x.cpu().detach().numpy()\ny = y.cpu().detach().numpy()","metadata":{"execution":{"iopub.status.busy":"2022-02-13T06:04:57.622409Z","iopub.execute_input":"2022-02-13T06:04:57.622942Z","iopub.status.idle":"2022-02-13T06:04:57.627131Z","shell.execute_reply.started":"2022-02-13T06:04:57.622903Z","shell.execute_reply":"2022-02-13T06:04:57.626437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x0 = x[18]\nx0 = np.transpose(x0,(1,2,0))\nplt.imshow(x0)","metadata":{"execution":{"iopub.status.busy":"2022-02-13T06:05:07.679830Z","iopub.execute_input":"2022-02-13T06:05:07.680144Z","iopub.status.idle":"2022-02-13T06:05:07.939308Z","shell.execute_reply.started":"2022-02-13T06:05:07.680106Z","shell.execute_reply":"2022-02-13T06:05:07.938519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y0 = y[18]\ny0 = np.transpose(y0,(1,2,0))\nplt.imshow(y0)","metadata":{"execution":{"iopub.status.busy":"2022-02-13T06:05:11.096443Z","iopub.execute_input":"2022-02-13T06:05:11.097167Z","iopub.status.idle":"2022-02-13T06:05:11.284354Z","shell.execute_reply.started":"2022-02-13T06:05:11.097129Z","shell.execute_reply":"2022-02-13T06:05:11.283603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x0 = x[54]\nx0 = np.transpose(x0,(1,2,0))\nplt.imshow(x0)","metadata":{"execution":{"iopub.status.busy":"2022-02-13T06:07:53.376277Z","iopub.execute_input":"2022-02-13T06:07:53.376524Z","iopub.status.idle":"2022-02-13T06:07:53.560049Z","shell.execute_reply.started":"2022-02-13T06:07:53.376497Z","shell.execute_reply":"2022-02-13T06:07:53.559377Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y0 = y[54]\ny0 = np.transpose(y0,(1,2,0))\nplt.imshow(y0)","metadata":{"execution":{"iopub.status.busy":"2022-02-13T06:07:56.645993Z","iopub.execute_input":"2022-02-13T06:07:56.646739Z","iopub.status.idle":"2022-02-13T06:07:56.823386Z","shell.execute_reply.started":"2022-02-13T06:07:56.646692Z","shell.execute_reply":"2022-02-13T06:07:56.822734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x = train_data[0].numpy()","metadata":{"execution":{"iopub.status.busy":"2022-02-10T04:46:48.24186Z","iopub.execute_input":"2022-02-10T04:46:48.242587Z","iopub.status.idle":"2022-02-10T04:46:48.286458Z","shell.execute_reply.started":"2022-02-10T04:46:48.242545Z","shell.execute_reply":"2022-02-10T04:46:48.2857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img1 =cv2.imread(\"../input/h-and-m-personalized-fashion-recommendations/images/010/0108775015.jpg\")\nimg = img1 / 256\nimg.shape","metadata":{"execution":{"iopub.status.busy":"2022-02-10T05:06:12.547286Z","iopub.execute_input":"2022-02-10T05:06:12.547557Z","iopub.status.idle":"2022-02-10T05:06:12.599178Z","shell.execute_reply.started":"2022-02-10T05:06:12.547526Z","shell.execute_reply":"2022-02-10T05:06:12.598281Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.transpose(img,(2,1,0)).shape","metadata":{"execution":{"iopub.status.busy":"2022-02-10T05:07:03.291365Z","iopub.execute_input":"2022-02-10T05:07:03.291723Z","iopub.status.idle":"2022-02-10T05:07:03.298183Z","shell.execute_reply.started":"2022-02-10T05:07:03.291682Z","shell.execute_reply":"2022-02-10T05:07:03.297431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img2 = train_data[1523].numpy()\nimg2 = np.transpose(img2, (1,2,0))\nplt.imshow(img2)","metadata":{"execution":{"iopub.status.busy":"2022-02-10T05:13:34.229436Z","iopub.execute_input":"2022-02-10T05:13:34.229999Z","iopub.status.idle":"2022-02-10T05:13:34.496243Z","shell.execute_reply.started":"2022-02-10T05:13:34.229958Z","shell.execute_reply":"2022-02-10T05:13:34.495528Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = cv2.cvtColor(x, cv2.COLOR_BGR2RGB)\nplt.imshow(img*256)","metadata":{"execution":{"iopub.status.busy":"2022-02-10T04:53:04.062461Z","iopub.execute_input":"2022-02-10T04:53:04.062734Z","iopub.status.idle":"2022-02-10T04:53:04.247461Z","shell.execute_reply.started":"2022-02-10T04:53:04.062706Z","shell.execute_reply":"2022-02-10T04:53:04.246783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(np.transpose(x, (1,2,0)))","metadata":{"execution":{"iopub.status.busy":"2022-02-10T04:45:48.191545Z","iopub.execute_input":"2022-02-10T04:45:48.191817Z","iopub.status.idle":"2022-02-10T04:45:48.388988Z","shell.execute_reply.started":"2022-02-10T04:45:48.191788Z","shell.execute_reply":"2022-02-10T04:45:48.388292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(np.transpose(x[12], (1,2,0)))","metadata":{"execution":{"iopub.status.busy":"2022-02-10T04:38:11.350484Z","iopub.execute_input":"2022-02-10T04:38:11.351011Z","iopub.status.idle":"2022-02-10T04:38:11.534048Z","shell.execute_reply.started":"2022-02-10T04:38:11.350972Z","shell.execute_reply":"2022-02-10T04:38:11.533354Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = np.transpose(result[3], (2,1,0))\nplt.imshow(img)","metadata":{"execution":{"iopub.status.busy":"2022-02-10T04:34:13.566035Z","iopub.execute_input":"2022-02-10T04:34:13.566764Z","iopub.status.idle":"2022-02-10T04:34:13.763589Z","shell.execute_reply.started":"2022-02-10T04:34:13.566725Z","shell.execute_reply":"2022-02-10T04:34:13.762917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_img = cv2.imread('../input/h-and-m-personalized-fashion-recommendations/images/010/0108775015.jpg')\nsample_img.shape","metadata":{"execution":{"iopub.status.busy":"2022-02-10T04:31:24.002337Z","iopub.execute_input":"2022-02-10T04:31:24.002593Z","iopub.status.idle":"2022-02-10T04:31:24.033638Z","shell.execute_reply.started":"2022-02-10T04:31:24.002564Z","shell.execute_reply":"2022-02-10T04:31:24.032897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(result[0])","metadata":{"execution":{"iopub.status.busy":"2022-02-10T04:30:36.183811Z","iopub.execute_input":"2022-02-10T04:30:36.184101Z","iopub.status.idle":"2022-02-10T04:30:36.418487Z","shell.execute_reply.started":"2022-02-10T04:30:36.184066Z","shell.execute_reply":"2022-02-10T04:30:36.416875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"z = z.cpu().detach().numpy()","metadata":{"execution":{"iopub.status.busy":"2022-02-10T04:23:06.656021Z","iopub.execute_input":"2022-02-10T04:23:06.656669Z","iopub.status.idle":"2022-02-10T04:23:06.674125Z","shell.execute_reply.started":"2022-02-10T04:23:06.656635Z","shell.execute_reply":"2022-02-10T04:23:06.673187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.neighbors import NearestNeighbors\nextractor = NearestNeighbors(metric='euclidean', n_neighbors=4)\nz[1:].shape","metadata":{"execution":{"iopub.status.busy":"2022-02-10T04:23:41.08397Z","iopub.execute_input":"2022-02-10T04:23:41.084538Z","iopub.status.idle":"2022-02-10T04:23:41.090379Z","shell.execute_reply.started":"2022-02-10T04:23:41.084498Z","shell.execute_reply":"2022-02-10T04:23:41.089683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}