{"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":"import numpy as np\nimport pandas as pd\n\n\ndf = pd.read_csv(\"../input/h-and-m-personalized-fashion-recommendations/transactions_train.csv\", dtype={\"article_id\": str})\nprint(df.shape)\ndf.head()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-22T22:17:28.880932Z","iopub.execute_input":"2022-08-22T22:17:28.881276Z","iopub.status.idle":"2022-08-22T22:18:25.986360Z","shell.execute_reply.started":"2022-08-22T22:17:28.881174Z","shell.execute_reply":"2022-08-22T22:18:25.985598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df[\"t_dat\"] = pd.to_datetime(df[\"t_dat\"])\ndf[\"t_dat\"].max()","metadata":{"execution":{"iopub.status.busy":"2022-08-22T22:18:25.988211Z","iopub.execute_input":"2022-08-22T22:18:25.988529Z","iopub.status.idle":"2022-08-22T22:18:30.775667Z","shell.execute_reply.started":"2022-08-22T22:18:25.988491Z","shell.execute_reply":"2022-08-22T22:18:30.774880Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"active_articles = df.groupby(\"article_id\")[\"t_dat\"].max().reset_index()\nactive_articles = active_articles[active_articles[\"t_dat\"] >= \"2019-09-01\"].reset_index()\nactive_articles.shape","metadata":{"execution":{"iopub.status.busy":"2022-08-22T22:18:30.779260Z","iopub.execute_input":"2022-08-22T22:18:30.779972Z","iopub.status.idle":"2022-08-22T22:18:35.487858Z","shell.execute_reply.started":"2022-08-22T22:18:30.779934Z","shell.execute_reply":"2022-08-22T22:18:35.487138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = df[df[\"article_id\"].isin(active_articles[\"article_id\"])].reset_index(drop=True)\ndf.shape","metadata":{"execution":{"iopub.status.busy":"2022-08-22T22:18:35.489481Z","iopub.execute_input":"2022-08-22T22:18:35.489738Z","iopub.status.idle":"2022-08-22T22:18:42.279658Z","shell.execute_reply.started":"2022-08-22T22:18:35.489701Z","shell.execute_reply":"2022-08-22T22:18:42.278951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df[\"week\"] = (df[\"t_dat\"].max() - df[\"t_dat\"]).dt.days // 7\ndf[\"week\"].value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-08-22T22:18:42.281884Z","iopub.execute_input":"2022-08-22T22:18:42.282137Z","iopub.status.idle":"2022-08-22T22:18:43.698978Z","shell.execute_reply.started":"2022-08-22T22:18:42.282103Z","shell.execute_reply":"2022-08-22T22:18:43.698137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.preprocessing import LabelEncoder\n\n\narticle_ids = np.concatenate([[\"placeholder\"], np.unique(df[\"article_id\"].values)])\n\nle_article = LabelEncoder()\nle_article.fit(article_ids)\ndf[\"article_id\"] = le_article.transform(df[\"article_id\"])","metadata":{"execution":{"iopub.status.busy":"2022-08-22T22:18:43.700310Z","iopub.execute_input":"2022-08-22T22:18:43.700560Z","iopub.status.idle":"2022-08-22T22:19:30.120035Z","shell.execute_reply.started":"2022-08-22T22:18:43.700526Z","shell.execute_reply":"2022-08-22T22:19:30.119095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"WEEK_HIST_MAX = 5\n\ndef create_dataset(df, week):\n    hist_df = df[(df[\"week\"] > week) & (df[\"week\"] <= week + WEEK_HIST_MAX)]\n    hist_df = hist_df.groupby(\"customer_id\").agg({\"article_id\": list, \"week\": list}).reset_index()\n    hist_df.rename(columns={\"week\": 'week_history'}, inplace=True)\n    \n    target_df = df[df[\"week\"] == week]\n    target_df = target_df.groupby(\"customer_id\").agg({\"article_id\": list}).reset_index()\n    target_df.rename(columns={\"article_id\": \"target\"}, inplace=True)\n    target_df[\"week\"] = week\n    \n    return target_df.merge(hist_df, on=\"customer_id\", how=\"left\")\n\nval_weeks = [0]\ntrain_weeks = [1, 2, 3, 4]\n\n\nval_df = pd.concat([create_dataset(df, w) for w in val_weeks]).reset_index(drop=True)\ntrain_df = pd.concat([create_dataset(df, w) for w in train_weeks]).reset_index(drop=True)\ntrain_df.shape, val_df.shape","metadata":{"execution":{"iopub.status.busy":"2022-08-22T22:19:30.121558Z","iopub.execute_input":"2022-08-22T22:19:30.121835Z","iopub.status.idle":"2022-08-22T22:20:01.129166Z","shell.execute_reply.started":"2022-08-22T22:19:30.121788Z","shell.execute_reply":"2022-08-22T22:20:01.128507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def adjust_lr(optimizer, epoch):\n    if epoch < 1:\n        lr = 5e-5\n    elif epoch < 6:\n        lr = 1e-3\n    elif epoch < 9:\n        lr = 1e-4\n    else:\n        lr = 1e-5\n\n    for p in optimizer.param_groups:\n        p['lr'] = lr\n    return lr\n    \ndef get_optimizer(net):\n    optimizer = torch.optim.Adam(filter(lambda p: p.requires_grad, net.parameters()), lr=3e-4, betas=(0.9, 0.999),\n                                 eps=1e-08)\n    return optimizer","metadata":{"execution":{"iopub.status.busy":"2022-08-22T22:34:19.774683Z","iopub.execute_input":"2022-08-22T22:34:19.774973Z","iopub.status.idle":"2022-08-22T22:34:19.781640Z","shell.execute_reply.started":"2022-08-22T22:34:19.774942Z","shell.execute_reply":"2022-08-22T22:34:19.780750Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torch\nfrom tqdm import tqdm\n","metadata":{"execution":{"iopub.status.busy":"2022-08-22T22:34:21.581378Z","iopub.execute_input":"2022-08-22T22:34:21.582042Z","iopub.status.idle":"2022-08-22T22:34:21.586989Z","shell.execute_reply.started":"2022-08-22T22:34:21.582004Z","shell.execute_reply":"2022-08-22T22:34:21.585978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TreeStructure(nn.Module):\n    def __init__(self, middle_index, item_index,layer_top_emb,layer_bottom_emb,first_train):\n        super(TreeStructure, self).__init__()\n\n        # Parameters\n        self.first_train = first_train\n        self.ntokens = 72582#the number of ouput.(72582)\n        self.nhid = 512#dimension: the same length of customer dimension.(512)\n\n        self.ntokens_per_class = 20#how many children one intermidiate node.(20)\n\n        self.nclasses = int(np.ceil(self.ntokens * 1. / self.ntokens_per_class))#intermidiate nodes.(3630)\n        self.ntokens_actual = self.nclasses * self.ntokens_per_class#72600\n        if self.first_train:\n            self.layer_top_emb = nn.Parameter(torch.FloatTensor(self.nclasses,self.nhid), requires_grad=True)\n            self.layer_bottom_emb = nn.Parameter(torch.FloatTensor(self.ntokens_actual, self.nhid), requires_grad=True)\n            self.init_weights()\n            #for K-means to cluster the embedding.(Initialization)\n            self.middle_index = np.arange(self.nclasses).tolist()\n            self.item_index = np.arange(self.ntokens_actual).tolist()\n        else:\n            #(Inherit from the previous K-means clustering)\n            self.middle_index = middle_index.tolist()\n            self.item_index = item_index.tolist()\n            self.layer_top_emb = nn.Parameter(layer_top_emb, requires_grad=True)\n            self.layer_bottom_emb = nn.Parameter(layer_bottom_emb, requires_grad=True)\n            \n\n    def init_weights(self):\n\n        initrange = 0.1\n        self.layer_top_emb.data.uniform_(-initrange, initrange)\n        self.layer_bottom_emb.data.uniform_(-initrange, initrange)\n\n\n    def forward(self, purchase_hist_npos):\n        #leaf index \n        hist = purchase_hist_npos\n        \n        #nonleaf index\n        parent_index = (hist/ self.ntokens_per_class).long()#the position after clustering \n\n        #leaf embedding\n        positive_leaf_emb = self.layer_bottom_emb[hist]#positive 1###[256, 512]\n        negative_leaf_sample = torch.LongTensor(np.random.choice(72600, positive_leaf_emb.shape[0]))###[256] \n        negative_leaf_emb = self.layer_bottom_emb[negative_leaf_sample]#negative 1###[256, 512]\n        #nonleaf embedding\n        positive_nonleaf_emb = self.layer_top_emb[parent_index]#positive 2###[256, 512]\n        negative_nonleaf_sample = torch.LongTensor(np.random.choice(3630, positive_leaf_emb.shape[0]))\n        negative_nonleaf_emb = self.layer_top_emb[negative_nonleaf_sample]#negative 2\n        \n        \n        return [positive_leaf_emb,negative_leaf_emb,positive_nonleaf_emb,negative_nonleaf_emb]\n","metadata":{"execution":{"iopub.status.busy":"2022-08-22T22:34:23.511899Z","iopub.execute_input":"2022-08-22T22:34:23.512333Z","iopub.status.idle":"2022-08-22T22:34:23.523395Z","shell.execute_reply.started":"2022-08-22T22:34:23.512285Z","shell.execute_reply":"2022-08-22T22:34:23.522622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HMModel(nn.Module):\n    def __init__(self, article_shape,first_train,middle_index, item_index,layer_top_emb,layer_bottom_emb,pre_emb):\n        super(HMModel, self).__init__()\n        \n        self.first_train = first_train\n        if self.first_train:\n            self.article_emb = torch.nn.Embedding(article_shape[0], embedding_dim=article_shape[1])\n            middle_index = torch.ones(1)\n            item_index = torch.ones(1)\n            layer_top_emb = torch.ones(1)\n            layer_bottom_emb = torch.ones(1)\n        else:\n            self.article_emb = torch.nn.Embedding.from_pretrained(torch.from_numpy(pre_emb).float())\n            self.middle_index = middle_index\n            self.item_index = item_index\n            self.layer_top_emb = layer_top_emb\n            self.layer_bottom_emb = layer_bottom_emb\n            \n        self.Tree = TreeStructure(middle_index, item_index,layer_top_emb,layer_bottom_emb,first_train=self.first_train)\n    def forward(self, inputs):\n        article_hist, week_hist, purchase_hist_npos = inputs[0], inputs[1], inputs[2]\n        x = self.article_emb(article_hist)\n        x = F.normalize(x, dim=2)###[256, 16, 512]\n        \n        x, indices = x.max(axis=1)##customer_emb[256,512]\n        \n        customer_emb = x\n        \n        global is_test\n        \n        if is_test:\n            \n            return customer_emb\n\n        #print('0',purchase_hist_item,purchase_hist_item.shape)\n        \n        [p1,n1,p2,n2] = self.Tree(purchase_hist_npos)#get four logits for 2 positive and 2 negative samples\n        \n        p1_dot = torch.mul(x,p1).sum(dim=1).unsqueeze(0)\n        n1_dot = torch.mul(x,n1).sum(dim=1).unsqueeze(0)\n        p2_dot = torch.mul(x,p2).sum(dim=1).unsqueeze(0)\n        n2_dot = torch.mul(x,n2).sum(dim=1).unsqueeze(0)\n        \n        \n        logits = torch.cat((p1_dot,n1_dot,p2_dot,n2_dot),0).T\n        \n        return logits\n\n#Train the model at the first time.\nmiddle_index = torch.ones(1)\nitem_index = torch.ones(1)\nlayer_top_emb = torch.ones(1)\nlayer_bottom_emb = torch.ones(1)\nfirst_train = True\nglobal first\nglobal is_test\nis_test = False\nfirst = True\narticle_emb = torch.ones(1)\n\nmodel = HMModel((len(le_article.classes_), 512),first_train,middle_index, item_index,layer_top_emb,layer_bottom_emb,article_emb)\nmodel = model.cuda()","metadata":{"execution":{"iopub.status.busy":"2022-08-22T22:34:26.772113Z","iopub.execute_input":"2022-08-22T22:34:26.772927Z","iopub.status.idle":"2022-08-22T22:34:27.772826Z","shell.execute_reply.started":"2022-08-22T22:34:26.772887Z","shell.execute_reply":"2022-08-22T22:34:27.771867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"global find_index\nglobal item_index\nglobal first_time\nfirst_time = True\nfind_index = torch.ones(1)\nitem_index = torch.ones(1)","metadata":{"execution":{"iopub.status.busy":"2022-08-22T22:34:29.989140Z","iopub.execute_input":"2022-08-22T22:34:29.989739Z","iopub.status.idle":"2022-08-22T22:34:29.994906Z","shell.execute_reply.started":"2022-08-22T22:34:29.989701Z","shell.execute_reply":"2022-08-22T22:34:29.994178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from scipy.spatial import distance\ndef compare_center_distance(center1,center2):\n    finalcenter2 = []\n    for c1 in center1:\n        for ini in range(len(center2)):#initialize.\n            if ini not in finalcenter2:\n                dis = distance.euclidean(ini,c1)\n                finalcenter2.append(ini)\n\n        for c in range(len(center2)):#compare to get the minimum distance.\n            if c in finalcenter2: continue\n            dis = min(dis,distance.euclidean(c1,c))\n            finalcenter2[-1] = c\n    return finalcenter2","metadata":{"execution":{"iopub.status.busy":"2022-08-22T22:23:22.071976Z","iopub.execute_input":"2022-08-22T22:23:22.072757Z","iopub.status.idle":"2022-08-22T22:23:22.078703Z","shell.execute_reply.started":"2022-08-22T22:23:22.072718Z","shell.execute_reply":"2022-08-22T22:23:22.077715Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.cluster import KMeans\ndef k_means(model):\n    #Find the current embedding of the tree.\n    cur_intermidiate_emb = model.Tree.layer_top_emb\n    cur_bottom_emb = model.Tree.layer_bottom_emb\n    article_emb = model.article_emb.weight.detach().cpu().numpy()\n    \n    #Find the current index of items in the tree.\n    middle_index = np.array(model.Tree.middle_index)\n    global item_index\n    item_index = np.array(model.Tree.item_index)\n    \n    #Use embedding to cluster(Get the index after K-means)\n    kmeans1 = KMeans(n_clusters=10, random_state=0).fit(cur_intermidiate_emb.cpu().detach().numpy())\n    global find_index\n    find_index = np.argsort(kmeans1.labels_)\n    middle_cluster = middle_index[find_index]\n    middle_centers = kmeans1.cluster_centers_\n    print('middle nodes clustering by K-means finished.')\n\n    kmeans2 = KMeans(n_clusters=10, random_state=0).fit(cur_bottom_emb.cpu().detach().numpy())\n    leaf_centers = kmeans2.cluster_centers_\n    \n    \n    new_label_order = compare_center_distance(middle_centers,leaf_centers)#get new leaf label order.\n    pre_labels = kmeans2.labels_#transform into new label catogory.[2,2,1,0,0]-->[1,1,0,2,2]\n    for i in range(len(pre_labels)):\n        pre_labels[i] = new_label_order[pre_labels[i]]\n    \n    \n    item_cluster = item_index[np.argsort(pre_labels)]\n    print('leaf nodes clustering by K-means finished.')\n    \n    #Reconstruct embedding matrix using above index(reconstruct tree)\n    suc_intermidiate_emb = torch.from_numpy(cur_intermidiate_emb.cpu().detach().numpy()[middle_cluster])\n    suc_bottom_emb = torch.from_numpy(cur_bottom_emb.cpu().detach().numpy()[item_cluster])\n          \n           ###3630          ###72600      ###[3630, 512]       ###[72600, 512]\n    return middle_cluster, item_cluster, suc_intermidiate_emb, suc_bottom_emb, article_emb\n    ","metadata":{"execution":{"iopub.status.busy":"2022-08-22T22:41:30.957073Z","iopub.execute_input":"2022-08-22T22:41:30.957897Z","iopub.status.idle":"2022-08-22T22:41:30.967660Z","shell.execute_reply.started":"2022-08-22T22:41:30.957861Z","shell.execute_reply":"2022-08-22T22:41:30.966877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HMDataset(Dataset):\n    def __init__(self, df, seq_len, model,is_test=False):\n        self.df = df.reset_index(drop=True)\n        self.seq_len = seq_len\n        self.is_test = is_test\n        self.model = model\n    \n    def __len__(self):\n        return self.df.shape[0]\n    \n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n        \n        if self.is_test:\n            target = torch.zeros(2).float()\n        else:\n            if not row.target:\n                target = torch.tensor([0]).int()\n            else:\n                rand_target = np.random.choice(row.target,1)\n                target = torch.tensor(rand_target).squeeze().int()\n#             for t in row.target:\n#                 target[t] = 1.0\n#                 break\n            \n        article_hist = torch.zeros(self.seq_len).long()\n        week_hist = torch.ones(self.seq_len).float()\n        \n        \n        if isinstance(row.article_id, list):\n            if len(row.article_id) >= self.seq_len:\n                article_hist = torch.LongTensor(row.article_id[-self.seq_len:])\n                week_hist = (torch.LongTensor(row.week_history[-self.seq_len:]) - row.week)/WEEK_HIST_MAX/2\n            else:\n                article_hist[-len(row.article_id):] = torch.LongTensor(row.article_id)\n                week_hist[-len(row.article_id):] = (torch.LongTensor(row.week_history) - row.week)/WEEK_HIST_MAX/2\n        target = torch.tensor([1,0,1,0]).float()\n        \n        purchase_hist_item = article_hist[-1].numpy()\n    \n        tree_item_index = self.model.Tree.item_index\n        global first_time\n        global find_index\n        global item_index\n        if first_time:\n            purchase_hist_npos = torch.tensor(purchase_hist_item)\n        else:\n            purchase_hist_npos = torch.tensor(find_index[item_index.index(purchase_hist_item)])\n        \n\n        return article_hist, week_hist, purchase_hist_npos, target\n    \nHMDataset(val_df, 64,model)[100]","metadata":{"execution":{"iopub.status.busy":"2022-08-22T22:34:36.513306Z","iopub.execute_input":"2022-08-22T22:34:36.513748Z","iopub.status.idle":"2022-08-22T22:34:36.566521Z","shell.execute_reply.started":"2022-08-22T22:34:36.513710Z","shell.execute_reply":"2022-08-22T22:34:36.565727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\n\ndef calc_map(topk_preds, target_array, k=12):\n    metric = []\n    tp, fp = 0, 0\n    \n    for pred in topk_preds:\n        if target_array[pred]:\n            tp += 1\n            metric.append(tp/(tp + fp))\n        else:\n            fp += 1\n            \n    return np.sum(metric) / min(k, target_array.sum())\n\ndef read_data(data):\n    return tuple(d.cuda() for d in data[:-1]), data[-1].cuda()\n    #return tuple(d for d in data[:-1]), data[-1]\n\n\ndef validate(model, val_loader, k=12):\n    model.eval()\n    \n    tbar = tqdm(val_loader, file=sys.stdout)\n    \n    maps = []\n    \n    with torch.no_grad():\n        for idx, data in enumerate(tbar):\n            inputs, target = read_data(data)\n\n            logits = model(inputs)\n\n            _, indices = torch.topk(logits, k, dim=1)\n\n            indices = indices.detach().cpu().numpy()\n            target = target.detach().cpu().numpy()\n            \n            for i in range(indices.shape[0]):\n                maps.append(calc_map(indices[i], target[i]))\n        \n    \n    return np.mean(maps)\n\nSEQ_LEN = 16\n\nBS = 256\nNW = 8\n\nval_dataset = HMDataset(val_df, SEQ_LEN,model)\nval_loader = DataLoader(val_dataset, batch_size=BS, shuffle=False, num_workers=NW,\n                          pin_memory=False, drop_last=False)","metadata":{"execution":{"iopub.status.busy":"2022-08-22T22:34:42.511483Z","iopub.execute_input":"2022-08-22T22:34:42.512113Z","iopub.status.idle":"2022-08-22T22:34:42.526568Z","shell.execute_reply.started":"2022-08-22T22:34:42.512074Z","shell.execute_reply":"2022-08-22T22:34:42.525026Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Train and validate","metadata":{}},{"cell_type":"code","source":"def dice_loss(y_pred, y_true):\n    y_pred = y_pred.sigmoid()\n    intersect = (y_true*y_pred).sum(axis=1)\n    \n    return 1 - (intersect/(intersect + y_true.sum(axis=1) + y_pred.sum(axis=1))).mean()\n\n\ndef train(model, train_loader, val_loader, epochs):\n    np.random.seed(SEED)\n    \n    optimizer = get_optimizer(model)\n    scaler = torch.cuda.amp.GradScaler()\n\n    criterion = nn.BCEWithLogitsLoss()\n    \n    for e in range(epochs):\n        model.train()\n        tbar = tqdm(train_loader, file=sys.stdout)\n        \n        lr = adjust_lr(optimizer, e)\n        \n        loss_list = []\n\n        for idx, data in enumerate(tbar):\n            inputs, target = read_data(data)\n\n            optimizer.zero_grad()\n            \n            with torch.cuda.amp.autocast():\n                logits = model(inputs)\n                \n#                 print('test1',logits,logits.shape)\n#                 print('test2',target,target.shape)\n                \n                loss = criterion(logits, target.float())\n            #loss.backward()\n            scaler.scale(loss).backward()\n            #optimizer.step()\n            scaler.step(optimizer)\n            scaler.update()\n            \n            loss_list.append(loss.detach().cpu().item())\n            \n            avg_loss = np.round(100*np.mean(loss_list), 4)\n\n            tbar.set_description(f\"Epoch {e+1} Loss: {avg_loss} lr: {lr}\")\n            \n#         val_map = validate(model, val_loader)\n\n#         log_text = f\"Epoch {e+1}\\nTrain Loss: {avg_loss}\\nValidation MAP: {val_map}\\n\"\n            \n#         print(log_text)\n        \n        #logfile = open(f\"models/{MODEL_NAME}_{SEED}.txt\", 'a')\n        #logfile.write(log_text)\n        #logfile.close()\n    return model\n\n\nMODEL_NAME = \"exp001\"\nSEED = 0\n\ntrain_dataset = HMDataset(train_df, SEQ_LEN,model)\ntrain_loader = DataLoader(train_dataset, batch_size=BS, shuffle=True, num_workers=NW,\n                          pin_memory=False, drop_last=True)\n\n#first training(Initializing)\nmodel = train(model, train_loader, val_loader, epochs=10) ","metadata":{"execution":{"iopub.status.busy":"2022-08-22T22:35:05.045724Z","iopub.execute_input":"2022-08-22T22:35:05.045987Z","iopub.status.idle":"2022-08-22T22:40:08.736348Z","shell.execute_reply.started":"2022-08-22T22:35:05.045956Z","shell.execute_reply":"2022-08-22T22:40:08.733342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Train to reconstruct the interest tree many times after first training.\nglobal first\nfirst = False\ndef train_tree(model,train_loader,val_loader,epochs):\n    cluster = k_means(model)\n    first_train = False\n    middle_index = cluster[0]\n    item_index = cluster[1]\n    layer_top_emb = cluster[2]\n    layer_bottom_emb = cluster[3]\n    article_emb = cluster[4]\n    model = HMModel((len(le_article.classes_), 512),first_train,middle_index, item_index,layer_top_emb,layer_bottom_emb,article_emb)\n    model = model.cuda()\n    return train(model, train_loader, val_loader, epochs)\n\nepochs = 10\ntrain_tree_epochs = 3\nfor _ in range(train_tree_epochs):\n    train_dataset = HMDataset(train_df, SEQ_LEN,model)\n    train_loader = DataLoader(train_dataset, batch_size=BS, shuffle=True, num_workers=NW,\n                          pin_memory=False, drop_last=True)\n    model = train_tree(model,train_loader,val_loader,epochs)\n","metadata":{"execution":{"iopub.status.busy":"2022-08-22T22:41:38.724673Z","iopub.execute_input":"2022-08-22T22:41:38.724946Z","iopub.status.idle":"2022-08-22T22:53:48.042937Z","shell.execute_reply.started":"2022-08-22T22:41:38.724915Z","shell.execute_reply":"2022-08-22T22:53:48.041686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"test_df = pd.read_csv('../input/h-and-m-personalized-fashion-recommendations/sample_submission.csv').drop(\"prediction\", axis=1)\nprint(test_df.shape)\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-21T02:21:15.350618Z","iopub.execute_input":"2022-08-21T02:21:15.351264Z","iopub.status.idle":"2022-08-21T02:21:17.743337Z","shell.execute_reply.started":"2022-08-21T02:21:15.351227Z","shell.execute_reply":"2022-08-21T02:21:17.742546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_test_dataset(test_df):\n    week = -1\n    test_df[\"week\"] = week\n    \n    hist_df = df[(df[\"week\"] > week) & (df[\"week\"] <= week + WEEK_HIST_MAX)]\n    hist_df = hist_df.groupby(\"customer_id\").agg({\"article_id\": list, \"week\": list}).reset_index()\n    hist_df.rename(columns={\"week\": 'week_history'}, inplace=True)\n    \n    \n    return test_df.merge(hist_df, on=\"customer_id\", how=\"left\")\n\ntest_df = create_test_dataset(test_df)\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-21T02:21:19.997241Z","iopub.execute_input":"2022-08-21T02:21:19.997873Z","iopub.status.idle":"2022-08-21T02:21:26.328909Z","shell.execute_reply.started":"2022-08-21T02:21:19.997822Z","shell.execute_reply":"2022-08-21T02:21:26.328183Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df[\"article_id\"].isnull().mean()","metadata":{"execution":{"iopub.status.busy":"2022-08-21T01:45:34.600884Z","iopub.execute_input":"2022-08-21T01:45:34.601479Z","iopub.status.idle":"2022-08-21T01:45:34.663856Z","shell.execute_reply.started":"2022-08-21T01:45:34.601433Z","shell.execute_reply":"2022-08-21T01:45:34.662927Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"global is_test\nis_test = True","metadata":{"execution":{"iopub.status.busy":"2022-08-21T02:21:55.446716Z","iopub.execute_input":"2022-08-21T02:21:55.447171Z","iopub.status.idle":"2022-08-21T02:21:55.450877Z","shell.execute_reply.started":"2022-08-21T02:21:55.447135Z","shell.execute_reply":"2022-08-21T02:21:55.450011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ds = HMDataset(test_df, SEQ_LEN, model,is_test=True)\ntest_loader = DataLoader(test_ds, batch_size=BS, shuffle=False, num_workers=NW,\n                          pin_memory=False, drop_last=False)\n\n\ndef inference(model, loader, k=12):\n    model.eval()\n    \n    tbar = tqdm(loader, file=sys.stdout)\n    \n    preds = []\n    \n    with torch.no_grad():\n        for idx, data in enumerate(tbar):\n            inputs, target = read_data(data)\n\n            customer_emb = model(inputs)\n                \n            dot_top_layer = torch.matmul(customer_emb,model.Tree.layer_top_emb.T)\n            _ , indices = torch.topk(dot_top_layer, 12, dim=1)\n            indices = indices.detach().cpu().numpy()\n\n            search_item = np.array(model.Tree.item_index)\n            search_item = np.where(search_item >= 72582, 0, search_item)\n            bottom_layer = model.Tree.layer_bottom_emb ##[0*indice,0*indice+20]\n            \n            for i in range(len(indices)):\n                for j in range(12):\n                    dot = torch.matmul(customer_emb[i],bottom_layer[indices[i][j]*20:indices[i][j]*20+20].T)\n                    _ , item_pos = torch.topk(dot, 1, dim=0)\n                    item_pos = indices[i][j]*20 + item_pos\n                    indices[i][j] = search_item[item_pos]\n   \n            for i in range(len(indices)):\n                preds.append(\" \".join(list(le_article.inverse_transform(indices[i]))))\n                \n    return preds\n\n\ntest_df[\"prediction\"] = inference(model, test_loader)","metadata":{"execution":{"iopub.status.busy":"2022-08-21T02:23:09.807927Z","iopub.execute_input":"2022-08-21T02:23:09.808396Z","iopub.status.idle":"2022-08-21T02:23:12.887363Z","shell.execute_reply.started":"2022-08-21T02:23:09.808359Z","shell.execute_reply":"2022-08-21T02:23:12.886325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df.to_csv(\"submission.csv\", index=False, columns=[\"customer_id\", \"prediction\"])","metadata":{"execution":{"iopub.status.busy":"2022-03-18T13:20:11.209147Z","iopub.status.idle":"2022-03-18T13:20:11.209811Z","shell.execute_reply.started":"2022-03-18T13:20:11.209573Z","shell.execute_reply":"2022-03-18T13:20:11.209598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}