{"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":"Hello World! This was the first kaggle competition that I committed to and tried to built something. Although its bits and pieces and highly unorganized and the overall architecture did not work, I got to learn a lot and felt like sharing the notebook for anyone interested. \nIn the notebook, I have implemented a Heterogenous Graph Neural Network using dgl package to generate and dump customer and articles embedding. Initial idea was to use this semantically rich embeddings and train a sequence model with attention and context. To be specific using customer embedding as context and the list of articles purchased as sequence(https://www.kaggle.com/code/mayankk9/transformermod). \nThis notebook is complete with data loader, model training and final embeddings dump. The above refered notebook is the continuation to use this embedding for sequence modelling.\nI understand the architecture was too complicated. However, my goal was to learn to use these architecture on real data. \nSuggestion\n- The other notebook for sequence currently uses transformer, both encoder and decoder, from pytorch. The pytorch implementation has many issues and would suggest to use hugging face models for this purpose for any future work.\n- Also, instead of transformer a LSTM/GRU based model for sequence modelling with context could have worked better\n\nHope it helps and you like my work. Please upvote if you find it interesting.","metadata":{}},{"cell_type":"code","source":"\n\n! nvcc --version; python --version\n! pip install dgl-cu110 dglgo -f https://data.dgl.ai/wheels/repo.html\n\n","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:54:19.166456Z","iopub.execute_input":"2022-05-13T00:54:19.166755Z","iopub.status.idle":"2022-05-13T00:54:57.755497Z","shell.execute_reply.started":"2022-05-13T00:54:19.166676Z","shell.execute_reply":"2022-05-13T00:54:57.754651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\nimport pandas as pd\npd.set_option(\"display.max_columns\", 100)\nimport cudf\nimport numpy as np\nimport os\nimport dgl\nimport torch\nimport cupy as cp\nfrom cuml.preprocessing.LabelEncoder import LabelEncoder\nfrom cuml.preprocessing.TargetEncoder import TargetEncoder\nimport joblib\nimport tqdm","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:54:57.757408Z","iopub.execute_input":"2022-05-13T00:54:57.75767Z","iopub.status.idle":"2022-05-13T00:55:04.068212Z","shell.execute_reply.started":"2022-05-13T00:54:57.757636Z","shell.execute_reply":"2022-05-13T00:55:04.063775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ! zip -r d.zip ../input/h-and-m-personalized-fashion-recommendations/transactions_train.csv","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:55:04.072175Z","iopub.execute_input":"2022-05-13T00:55:04.072444Z","iopub.status.idle":"2022-05-13T00:55:04.08285Z","shell.execute_reply.started":"2022-05-13T00:55:04.072408Z","shell.execute_reply":"2022-05-13T00:55:04.081832Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\nfeat = {'customers':{'cat':['club_member_status', 'fashion_news_frequency', 'postal_code'],\n       'cont':['age']}, \n        'articles': {'cat': ['product_code', 'product_type_no', 'product_group_name',\n              'graphical_appearance_no', 'colour_group_code','perceived_colour_master_id' , 'garment_group_no',\n              'section_no', 'index_group_no'],\n        'cont':[]}\n       }\n\n","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:55:04.091147Z","iopub.execute_input":"2022-05-13T00:55:04.09134Z","iopub.status.idle":"2022-05-13T00:55:04.967067Z","shell.execute_reply.started":"2022-05-13T00:55:04.091316Z","shell.execute_reply":"2022-05-13T00:55:04.966119Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import dgl\nfrom dgl.data import DGLDataset\nimport torch\nimport os\nimport numpy as np\nimport cudf\nfrom sklearn.preprocessing import MinMaxScaler","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:55:04.973385Z","iopub.execute_input":"2022-05-13T00:55:04.973896Z","iopub.status.idle":"2022-05-13T00:55:05.073052Z","shell.execute_reply.started":"2022-05-13T00:55:04.973853Z","shell.execute_reply":"2022-05-13T00:55:05.071137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install wandb -qqq\nimport wandb\nimport os\nos.environ[\"WANDB_API_KEY\"] = \"7e40122ef4016b01cdc4a60bcd7ee7a251b020a7\"\nwandb.init(project=\"HnmRGCNv2\", resume=True)","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:55:05.074224Z","iopub.execute_input":"2022-05-13T00:55:05.075143Z","iopub.status.idle":"2022-05-13T00:55:22.183968Z","shell.execute_reply.started":"2022-05-13T00:55:05.075105Z","shell.execute_reply":"2022-05-13T00:55:22.183153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import dgl\nfrom dgl.data import DGLDataset\nimport torch\nimport os\nimport numpy as np\nimport cudf\nfrom sklearn.preprocessing import MinMaxScaler\n\nclass HandMData(DGLDataset):\n    def __init__(self):\n        super().__init__(name='HandMData')\n\n    def process(self):\n        \n        base = \"../input/h-and-m-personalized-fashion-recommendations/\"\n        \n        nodes_articles = cudf.read_csv(f'./{base}/articles.csv')\n        nodes_articles['article_id'] = nodes_articles['article_id'].astype('int32')\n        le_a = LabelEncoder(handle_unknown='ignore')\n        le_a.fit(nodes_articles['article_id'])\n        \n        import pickle\n        import joblib\n\n        joblib.dump(le_a, 'label_encoder.joblib')\n#         label_encoder = joblib.load('label_encoder.joblib')\n        with open(\"le_a.pkl\", \"wb\") as fp:\n            pickle.dump(le_a, fp)\n        wandb.save(f\"le_a.pkl\")\n        nodes_articles['article_id'] = le_a.transform(nodes_articles['article_id'])\n        del le_a\n        nodes_articles = nodes_articles.sort_values('article_id')\n        \n        for col in feat['articles']['cat']:\n            le_a = LabelEncoder(handle_unknown='ignore')\n            le_a.fit(nodes_articles[col])\n#             print(le_a.classes_)\n            import pickle\n            with open(f\"le_{col}.pkl\", \"wb\") as fp:\n                pickle.dump(le_a, fp)\n            nodes_articles[col] = le_a.transform(nodes_articles[col])\n            del le_a\n        \n#         le_ap = LabelEncoder(handle_unknown='ignore')\n#         le_ap.fit(nodes_articles['product_code'].fillna('PC'))\n#         nodes_articles['product_code'] = le_ap.transform(nodes_articles['product_code'])\n#         with open(\"le_ap.pkl\", \"wb\") as fp:\n#             pickle.dump(le_ap, fp)\n#         del le_ap\n#         del nodes_articles\n        \n        nodes_customers = cudf.read_csv(f'./{base}/customers.csv')\n        nodes_customers['customer_id'] = nodes_customers['customer_id'].str[-16:].str.hex_to_int().astype('int64')\n        \n        le_c = LabelEncoder(handle_unknown='ignore')\n        le_c.fit(nodes_customers['customer_id'])\n        nodes_customers['customer_id'] = le_c.transform(nodes_customers['customer_id'])\n        import pickle\n        with open(\"le_c.pkl\", \"wb\") as fp:\n            pickle.dump(le_c, fp)\n        wandb.save(f\"le_c.pkl\")\n        del le_c\n        nodes_customers = nodes_customers.sort_values('customer_id')\n        # toodo include age at end MinMax Scaler\n        for col in feat['customers']['cat']:\n            le_a = LabelEncoder(handle_unknown='ignore')\n            le_a.fit(nodes_customers[col])\n#             print(le_a.classes_)\n            import pickle\n            with open(f\"le_{col}.pkl\", \"wb\") as fp:\n                pickle.dump(le_a, fp)\n            wandb.save(f\"le_{col}.pkl\")\n            nodes_customers[col] = le_a.transform(nodes_customers[col])\n            del le_a\n        \n        scaler = MinMaxScaler()\n        scaler.fit_transform(np.asarray(nodes_customers['age'].to_arrow()).reshape(-1,1))\n        with open(f\"MMxS_age.pkl\", \"wb\") as fp:\n            pickle.dump(scaler, fp)\n        wandb.save(f\"MMxS_age.pkl\")\n        transaction_df = cudf.read_csv(f'{base}/transactions_train.csv')\n        transaction_df['customer_id'] = transaction_df['customer_id'].str[-16:].str.hex_to_int().astype('int64')\n        transaction_df['article_id'] = transaction_df['article_id'].astype('int32')\n        \n        with open(\"le_c.pkl\", \"rb\") as fp:\n            le_c = pickle.load(fp)\n        \n        transaction_df['customer_id'] = le_c.transform(transaction_df['customer_id'])\n        del le_c #todo del unkown from here may be\n        \n        with open(\"le_a.pkl\", \"rb\") as fp:\n            le_a = pickle.load(fp)\n        transaction_df['article_id'] = le_a.transform(transaction_df['article_id'])\n        del le_a\n        \n        edges_data = transaction_df.groupby(['customer_id', 'article_id'])['t_dat'].count()\n        del transaction_df\n        \n        edges_data = edges_data.reset_index().rename(columns={'t_dat':'count'})\n        edge_features = torch.ones(edges_data.shape[0], dtype=torch.int32)#torch.from_numpy(cp.asnumpy(edges_data['count']))\n        #todo change edge_feature to make it regression link \n        graph_data = {\n        ('customers', 'buys', 'articles'): (torch.from_numpy(cp.asnumpy(edges_data['customer_id'].astype('int32'))),\n                                          torch.from_numpy(cp.asnumpy(edges_data['article_id'].astype('int32')))),\n        ('articles', 'boughtby', 'customers'): (torch.from_numpy(cp.asnumpy(edges_data['article_id'].astype('int32'))),\n                                             torch.from_numpy(cp.asnumpy(edges_data['customer_id'].astype('int32'))))}\n        \n        \n        del edges_data\n        # todo all doubles any help include feature for edges same torch float32 [[feat_x]]\n        self.g = dgl.heterograph(graph_data)\n        self.g.edges['buys'].data['u'] = edge_features\n        self.g.edges['boughtby'].data['u'] = edge_features\n        del edge_features\n        \n        cust_cols = feat['customers']['cat'] + feat['customers']['cont']\n        art_cols = feat['articles']['cat'] + feat['articles']['cont']\n        nodes_customers[cust_cols].fillna(0, inplace=True)\n        nodes_articles[art_cols].fillna(0, inplace=True)\n        \n        self.g.nodes['customers'].data['h'] = torch.from_numpy(nodes_customers[cust_cols].as_matrix().astype('float32'))\n        self.g.nodes['customers'].data['c'] = torch.ones(nodes_customers.shape[0])\n#         print(nodes_articles[art_cols].as_matrix())\n        self.g.nodes['articles'].data['h'] = torch.from_numpy(nodes_articles[art_cols].as_matrix().astype('float32'))[:-2]\n        self.g.nodes['articles'].data['c'] = torch.zeros(nodes_articles.shape[0])[:-2]\n        \n        del nodes_customers\n        del nodes_articles\n#         self.graph = dgl.graph((edges_src, edges_dst), num_nodes=nodes_data.shape[0])\n#         self.graph.ndata['feat'] = node_features\n#         self.graph.ndata['label'] = node_labels\n#         self.graph.edata['weight'] = edge_features\n\n        # If your dataset is a node classification dataset, you will need to assign\n        # masks indicating whether a node belongs to training, validation, and test set.\n        for sub_graph in ['customers', 'articles']:\n            n_nodes = self.g.num_nodes(f'{sub_graph}')\n            n_train = int(n_nodes * 0.7)\n            n_val = int(n_nodes * 0.15)\n            train_mask = torch.zeros(n_nodes, dtype=torch.bool)\n            val_mask = torch.zeros(n_nodes, dtype=torch.bool)\n            test_mask = torch.zeros(n_nodes, dtype=torch.bool)\n            train_mask[:n_train] = True\n            val_mask[n_train:n_train + n_val] = True\n            test_mask[n_train + n_val:] = True\n            self.g.nodes[f'{sub_graph}'].data['train_mask'] = train_mask\n            self.g.nodes[f'{sub_graph}'].data['val_mask'] = val_mask\n            self.g.nodes[f'{sub_graph}'].data['test_mask'] = test_mask\n        \n        for etype in self.g.etypes:\n            n_edges = self.g.edges(etype=etype, form='eid').shape[0]\n            n_train = int(n_edges * 0.7)\n            n_val = int(n_edges * 0.15)\n            train_mask = torch.zeros(n_edges, dtype=torch.bool)\n            val_mask = torch.zeros(n_edges, dtype=torch.bool)\n            test_mask = torch.zeros(n_edges, dtype=torch.bool)\n            train_mask[:n_train] = True\n            val_mask[n_train:n_train + n_val] = True\n            test_mask[n_train + n_val:] = True\n            self.g.edges[f'{etype}'].data['train_mask'] = train_mask\n            self.g.edges[f'{etype}'].data['val_mask'] = val_mask\n            self.g.edges[f'{etype}'].data['test_mask'] = test_mask\n        self.g = self.g.to(torch.device('cuda'))\n    def __getitem__(self, i):\n        return self.g\n\n    def __len__(self):\n        return 1","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:55:22.188709Z","iopub.execute_input":"2022-05-13T00:55:22.191025Z","iopub.status.idle":"2022-05-13T00:55:22.249893Z","shell.execute_reply.started":"2022-05-13T00:55:22.19098Z","shell.execute_reply":"2022-05-13T00:55:22.249117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dataset = HandMData()\ngraph = HandMData()[0]\nprint(graph)","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:55:22.254182Z","iopub.execute_input":"2022-05-13T00:55:22.256144Z","iopub.status.idle":"2022-05-13T00:56:14.085934Z","shell.execute_reply.started":"2022-05-13T00:55:22.256107Z","shell.execute_reply":"2022-05-13T00:56:14.085248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # definite infoleak as other set truth can be used as ng \n# train_eid_dict = {etype: (graph.edges[etype].data['train_mask'] == True).nonzero(as_tuple=True)[0] for etype in graph.etypes}\n# val_eid_dict   = {etype: (graph.edges[etype].data['val_mask'] == True).nonzero(as_tuple=True)[0] for etype in graph.etypes}\n# test_eid_dict   = {etype: (graph.edges[etype].data['test_mask'] == True).nonzero(as_tuple=True)[0] for etype in graph.etypes}","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.089649Z","iopub.execute_input":"2022-05-13T00:56:14.091765Z","iopub.status.idle":"2022-05-13T00:56:14.097354Z","shell.execute_reply.started":"2022-05-13T00:56:14.091726Z","shell.execute_reply":"2022-05-13T00:56:14.096711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del graph","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.104403Z","iopub.execute_input":"2022-05-13T00:56:14.106512Z","iopub.status.idle":"2022-05-13T00:56:14.151007Z","shell.execute_reply.started":"2022-05-13T00:56:14.106475Z","shell.execute_reply":"2022-05-13T00:56:14.15023Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\n# train_eid_dict = {k:v.to(torch.int32) for k,v in train_eid_dict.items()}\n# val_eid_dict = {k:v.to(torch.int32) for k,v in val_eid_dict.items()}\n# test_eid_dict = {k:v.to(torch.int32) for k,v in test_eid_dict.items()}\n\n","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.152185Z","iopub.execute_input":"2022-05-13T00:56:14.152651Z","iopub.status.idle":"2022-05-13T00:56:14.162302Z","shell.execute_reply.started":"2022-05-13T00:56:14.152616Z","shell.execute_reply":"2022-05-13T00:56:14.161579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# edge_sampler = dgl.dataloading.MultiLayerNeighborSampler([15, 10, 5])\n# negative_sampler=dgl.dataloading.negative_sampler.Uniform(5)\n\n# sampler = dgl.dataloading.as_edge_prediction_sampler(edge_sampler, exclude='reverse_types',\n#                                                      reverse_etypes={'buys': 'boughtby', 'boughtby': 'buys'},\n#                                                      negative_sampler= negative_sampler)\n                                                     \n# train_dataloader = dgl.dataloading.DataLoader(graph, train_eid_dict, sampler, \n#                                         batch_size=1024, shuffle=True, drop_last=False, num_workers=0)\n# val_dataloader = dgl.dataloading.DataLoader(graph, val_eid_dict, sampler, \n#                                         batch_size=1024, shuffle=True, drop_last=False, num_workers=0)\n# test_dataloader = dgl.dataloading.DataLoader(graph, test_eid_dict, sampler, \n#                                         batch_size=1024, shuffle=True, drop_last=False, num_workers=0)","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.163468Z","iopub.execute_input":"2022-05-13T00:56:14.163932Z","iopub.status.idle":"2022-05-13T00:56:14.182863Z","shell.execute_reply.started":"2022-05-13T00:56:14.163897Z","shell.execute_reply":"2022-05-13T00:56:14.18216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn as nn\nimport torch.nn.functional as F\nimport dgl.nn.pytorch as dglnn\nimport torch as th","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.184309Z","iopub.execute_input":"2022-05-13T00:56:14.184824Z","iopub.status.idle":"2022-05-13T00:56:14.21184Z","shell.execute_reply.started":"2022-05-13T00:56:14.184785Z","shell.execute_reply":"2022-05-13T00:56:14.211127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch.nn as nn\n# import torch.nn.functional as F\n# import dgl.nn.pytorch as dglnn\n# import torch as th\n\n# class StochasticThreeLayerRGCN(nn.Module):\n#     def __init__(self, in_feat_dict, hidden_feat, out_feat, canonical_etypes):\n#         super().__init__()\n\n#         self.conv1 = dglnn.HeteroGraphConv({\n#                 # you can change in_feat to something like in_feat[ntype] here\n#                 etype : dglnn.GraphConv(in_feat_dict[utype], hidden_feat, norm='right')\n#                 for utype, etype, vtype in canonical_etypes\n#                 })\n#         self.conv2 = dglnn.HeteroGraphConv({\n#                 etype : dglnn.GraphConv(hidden_feat, out_feat, norm='right')\n#                 for _, etype, _ in canonical_etypes\n#                 })\n#         self.conv3 = dglnn.HeteroGraphConv({\n#                 etype : dglnn.GraphConv(out_feat, out_feat, norm='right')\n#                 for _, etype, _ in canonical_etypes\n#                 })\n#     def forward(self, blocks, inputs):\n\n#         x1 = self.conv1(blocks[0], inputs)\n#         x2 = self.conv2(blocks[1], x1)\n#         x3 = self.conv3(blocks[2], x2)\n#         return x3\n\n# class HeteroScorePredictor(nn.Module):        \n#     def forward(self, edge_subgraph, x):\n#         with edge_subgraph.local_scope():\n# #             print(edge_subgraph.ndata)\n#             edge_subgraph.ndata['h'] = x\n#             for etype in edge_subgraph.canonical_etypes:\n#                 edge_subgraph.apply_edges(\n#                     dgl.function.u_dot_v('h', 'h', 'score'), etype=etype)\n#             return edge_subgraph.edata['score']\n\n# class Model(nn.Module):\n#     def __init__(self, in_features, hidden_features, out_features,canonical_etypes):\n#         super().__init__()\n#         self.rgcn = StochasticThreeLayerRGCN(in_features, hidden_features, out_features, canonical_etypes)\n#         self.pred = HeteroScorePredictor()\n    \n#     def _fetch_embLayer(self):\n#         return self.rgcn\n    \n#     def forward(self, graphs, positive_graph, negative_graph, blocks, x):\n#         x = self.rgcn(blocks, x)\n#         pos_score = self.pred(positive_graph, x)\n#         neg_score = self.pred(negative_graph, x)\n#         return x, pos_score, neg_score","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.215533Z","iopub.execute_input":"2022-05-13T00:56:14.215978Z","iopub.status.idle":"2022-05-13T00:56:14.231003Z","shell.execute_reply.started":"2022-05-13T00:56:14.215944Z","shell.execute_reply":"2022-05-13T00:56:14.230398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def compute_loss(pos_score, neg_score, canonical_etypes):\n#     # Margin loss\n#     all_losses = []\n#     for given_type in canonical_etypes:\n#         n_edges = pos_score[given_type].shape[0]\n#         if n_edges == 0:\n#             continue\n#         all_losses.append((1 - neg_score[given_type].view(n_edges, -1) + pos_score[given_type].unsqueeze(1)).clamp(min=0).mean())\n#     return torch.stack(all_losses, dim=0).mean()\n\n# from sklearn.metrics import roc_auc_score\n# import tensorflow as tf\n# def compute_acc(pos_score):\n#     crt = 0\n#     t = 0\n# #     ruc = 0 \n#     for k,v in pos_score.items():\n#         y_pred = pos_score[k].cpu().detach().numpy()\n#         m = tf.keras.metrics.BinaryAccuracy()\n#         y_true = torch.from_numpy(np.ones(y_pred.shape[0]))\n#         m.update_state(y_pred, y_true)\n#         crt += m.result().numpy()\n#         t += y_pred.shape[0]\n# #         ruc +=  roc_auc_score(y_true,y_pred)\n#     return crt/t","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.236801Z","iopub.execute_input":"2022-05-13T00:56:14.237293Z","iopub.status.idle":"2022-05-13T00:56:14.258296Z","shell.execute_reply.started":"2022-05-13T00:56:14.237247Z","shell.execute_reply":"2022-05-13T00:56:14.257624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def run_train_epoch(train_dataloader):\n#     ac_ = 0\n#     model.train()\n#     with tqdm.tqdm(train_dataloader) as tq:\n#         for step, (input_nodes, positive_graph, negative_graph, blocks) in enumerate(tq):\n            \n#             blocks = [b.to(torch.device('cuda')) for b in blocks]\n\n#             positive_graph = positive_graph.to(torch.device('cuda'))\n#             negative_graph = negative_graph.to(torch.device('cuda'))\n\n#             node_features = {'articles': blocks[0].srcdata['h']['articles'], \n#                              'customers': blocks[0].srcdata['h']['customers']}\n#             pos_score, neg_score = model(graph, positive_graph, negative_graph, blocks, node_features)\n#             loss = compute_loss(pos_score, neg_score, graph.canonical_etypes)\n#             optimizer.zero_grad()\n#             loss.backward()\n#             optimizer.step()\n#             tq.set_postfix({'loss': '%.03f' % loss.item()}, refresh=False)\n            \n#             if step % 1000==0:\n#                 torch.save(model.state_dict(), f'mb-{step}-train-{loss}-{best_model_path}')\n#                 lm = f'mb-{step}-train-{loss}-{best_model_path}'\n#                 wandb.save( lm, base_path=\"./\", policy=\"now\")\n#             if step % 5000==0:\n#                 for fnm in os.listdir():\n#                     if fnm.endswith('.pt'):\n#                         if fnm == lm:continue\n#                         os.remove(fnm)\n#             t = compute_acc(pos_score)          \n#             wandb.log({\"train-loss-mb\": loss})\n#             wandb.log({\"train-acc-mb\": t})\n#             ac_ += t\n            \n            \n        \n        \n#         wandb.log({\"train-acc-epochs\": ac_})\n#         print(\"Accuracy after epoch End on Train \", ac_)\n#         torch.save(model.state_dict(), f'epoch-train-{loss}-{best_model_path}')\n#         wandb.save( f'epoch-train-{loss}-{best_model_path}')\n        \n#         return loss, ac_","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.259206Z","iopub.execute_input":"2022-05-13T00:56:14.259428Z","iopub.status.idle":"2022-05-13T00:56:14.272772Z","shell.execute_reply.started":"2022-05-13T00:56:14.2594Z","shell.execute_reply":"2022-05-13T00:56:14.272144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ### function below ###\n# best_accuracy = 0\n# best_model_path = 'model.pt'\n\n\n# def run_eval_epoch(train_dataloader):\n#     global best_accuracy\n#     global best_model_path\n#     ac_ = 0\n#     model.eval()\n#     with tqdm.tqdm(train_dataloader) as tq:\n#         for step, (input_nodes, positive_graph, negative_graph, blocks) in enumerate(tq):\n#             blocks = [b.to(torch.device('cuda')) for b in blocks]\n\n#             positive_graph = positive_graph.to(torch.device('cuda'))\n#             negative_graph = negative_graph.to(torch.device('cuda'))\n\n#             node_features = {'articles': blocks[0].srcdata['h']['articles'], \n#                              'customers': blocks[0].srcdata['h']['customers']}\n#             with torch.no_grad():\n#                 pos_score, neg_score = model(graph, positive_graph, negative_graph, blocks, node_features)\n#                 ac = compute_acc(pos_score) # batch normalized already\n#                 if best_accuracy < ac:\n#                     best_accuracy = ac\n#                     torch.save(model.state_dict(), f'bm-{step}-val-{ac}-{best_model_path}')\n#                     wandb.save( f'bm-{step}-val-{ac}-{best_model_path}',base_path=\"./\", policy=\"now\")\n#                 loss = compute_loss(pos_score, neg_score, graph.canonical_etypes)\n#                 wandb.log({\"val-loss-mb\": loss})\n#                 tq.set_postfix({'loss': '%.03f' % loss.item()}, refresh=False)\n#                 t = compute_acc(pos_score)\n#                 ac_ += t\n#                 wandb.log({\"val-acc-mb\": t})\n                \n#         wandb.log({\"val-acc-epochs\": ac_})\n        \n#         print(\"Accuracy after epoch End on Val \", ac_)\n#         return loss, ac_","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.274136Z","iopub.execute_input":"2022-05-13T00:56:14.274693Z","iopub.status.idle":"2022-05-13T00:56:14.292948Z","shell.execute_reply.started":"2022-05-13T00:56:14.274656Z","shell.execute_reply":"2022-05-13T00:56:14.292295Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Code","metadata":{}},{"cell_type":"code","source":"# wandb.init(project=\"HnmRGCNv2\")\n# model = Model(in_features={'customers':4, 'articles':9},\n#               hidden_features=512, out_features=256, canonical_etypes=graph.canonical_etypes)\n# model = model.cuda()\n# optimizer = torch.optim.Adam(model.parameters())\n# import time\n# wandb.watch(model)\n          \n    \n# for epoch in range(4):\n#     train_loss, train_accuracy  = run_train_epoch(train_dataloader) # last batch loss and acc\n    \n#     eval_loss, eval_accuracy = run_eval_epoch(val_dataloader)    \n#     print(f'For epoch {epoch}: {eval_loss, eval_accuracy}')\n    \n    \n# torch.save(x, 'tensor.pt')  \n# torch.save(model.state_dict(), f'EOF-train-{best_model_path}')\n# wandb.save(f'EOF-train-{best_model_path}') ","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.294253Z","iopub.execute_input":"2022-05-13T00:56:14.294721Z","iopub.status.idle":"2022-05-13T00:56:14.304459Z","shell.execute_reply.started":"2022-05-13T00:56:14.294684Z","shell.execute_reply":"2022-05-13T00:56:14.302168Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load Model","metadata":{}},{"cell_type":"code","source":"# last_model = wandb.restore('EOF-train-model.pt', \n#                            run_path=\"mayankk-om-dev/HnmRGCNv2/1g06378a\")\n\n# # # use the \"name\" attribute of the returned object if your framework expects a filename, e.g. as in Keras\n# checkpoint = torch.load(last_model.name)\n\n# model = Model(in_features={'customers':4, 'articles':9},\n#               hidden_features=512, out_features=256, canonical_etypes=graph.canonical_etypes)\n# model = model.cuda()\n# model.load_state_dict(checkpoint)","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.305876Z","iopub.execute_input":"2022-05-13T00:56:14.309753Z","iopub.status.idle":"2022-05-13T00:56:14.320816Z","shell.execute_reply.started":"2022-05-13T00:56:14.30971Z","shell.execute_reply":"2022-05-13T00:56:14.319788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# graph","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.326315Z","iopub.execute_input":"2022-05-13T00:56:14.326783Z","iopub.status.idle":"2022-05-13T00:56:14.331447Z","shell.execute_reply.started":"2022-05-13T00:56:14.326745Z","shell.execute_reply":"2022-05-13T00:56:14.33076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# graph.number_of_nodes(\"articles\")","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.332753Z","iopub.execute_input":"2022-05-13T00:56:14.333226Z","iopub.status.idle":"2022-05-13T00:56:14.340611Z","shell.execute_reply.started":"2022-05-13T00:56:14.333189Z","shell.execute_reply":"2022-05-13T00:56:14.339944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dump Code Dont run load ","metadata":{}},{"cell_type":"code","source":"# def dump_embs(model, dataloader, embs_c, embs_a):#, embs_c, embs_a, dim_size):\n#     t0 = time.time()\n    \n#     model.eval()\n#     with tqdm.tqdm(dataloader) as tq:\n#         for step, (input_nodes, positive_graph, negative_graph, blocks) in enumerate(tq):\n#             blocks = [b.to(torch.device('cuda')) for b in blocks]\n\n#             positive_graph = positive_graph.to(torch.device('cuda'))\n            \n#             node_features = {'articles': blocks[0].srcdata['h']['articles'], \n#                              'customers': blocks[0].srcdata['h']['customers']}\n#             with torch.no_grad():\n#                 emb_feats, pos_score, neg_score = model(graph, positive_graph, negative_graph, blocks, node_features)\n#                 embs_c[blocks[-1].dstdata['_ID']['customers'].long()] = emb_feats['customers'].cpu()\n#                 embs_a[blocks[-1].dstdata['_ID']['articles'].long()] = emb_feats['articles'].cpu()\n# #             print(positive_graph.nodes['customers'].data['_ID'])\n# #             print(positive_graph.nodes['customers'].data['_ID'].shape)\n# #             print(blocks[-1])\n# #             print(type(blocks[-1]))\n# #             print(blocks[-1].dstdata['_ID']['customers'])\n# #             print(blocks[-1].dstdata['_ID']['customers'].shape)\n# #             print(type(x))\n# #             print(x['customers'].shape)\n# #             print(input_nodes[:5])\n# #             break\n#     print(f\"took {time.time()-t0} seconds\")\n#     return embs_c, embs_a       ","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.34231Z","iopub.execute_input":"2022-05-13T00:56:14.342822Z","iopub.status.idle":"2022-05-13T00:56:14.363755Z","shell.execute_reply.started":"2022-05-13T00:56:14.342783Z","shell.execute_reply":"2022-05-13T00:56:14.358171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# c = torch.empty(N, seq_len, 256)\n# emb = torch.randn(N, 256)\n# B = A.unsqueeze(1).repeat(1, K, 1)\n# import einops\n# einops.repeat(x, 'm n -> m k n', k=K)","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.367158Z","iopub.execute_input":"2022-05-13T00:56:14.367396Z","iopub.status.idle":"2022-05-13T00:56:14.380818Z","shell.execute_reply.started":"2022-05-13T00:56:14.367366Z","shell.execute_reply":"2022-05-13T00:56:14.379906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import tqdm\n# embs_c = torch.empty(graph.number_of_nodes(\"customers\"), 256)\n# embs_a = torch.empty(graph.number_of_nodes(\"articles\"), 256)\n\n# embs_c, embs_a  = dump_embs(model, train_dataloader, embs_c, embs_a)\n# embs_c, embs_a  = dump_embs(model, val_dataloader, embs_c, embs_a)\n# embs_c, embs_a  = dump_embs(model, test_dataloader, embs_c, embs_a)","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.381738Z","iopub.execute_input":"2022-05-13T00:56:14.381976Z","iopub.status.idle":"2022-05-13T00:56:14.3914Z","shell.execute_reply.started":"2022-05-13T00:56:14.381946Z","shell.execute_reply":"2022-05-13T00:56:14.386567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# wandb.init(project=\"HnmRGCNv2\")\n# torch.save(embs_c, 'customerGNNEmb.pt')\n# wandb.save('customerGNNEmb.pt') \n# torch.save(embs_a, 'articleGNNEmb.pt')  \n# wandb.save('articleGNNEmb.pt') ","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.392381Z","iopub.execute_input":"2022-05-13T00:56:14.392619Z","iopub.status.idle":"2022-05-13T00:56:14.408655Z","shell.execute_reply.started":"2022-05-13T00:56:14.392589Z","shell.execute_reply":"2022-05-13T00:56:14.403918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# embs_cD = wandb.restore('customerGNNEmb.pt', \n#                            run_path=\"mayankk-om-dev/HnmRGCNv2/2oooot9x\")\n\n# # # use the \"name\" attribute of the returned object if your framework expects a filename, e.g. as in Keras\n# embs_c = torch.load(embs_cD)\n# del embs_cD\n# embs_aD = wandb.restore('articleGNNEmb.pt', \n#                            run_path=\"mayankk-om-dev/HnmRGCNv2/2oooot9x\")\n# embs_a = torch.load(embs_aD)\n# del embs_aD","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.409846Z","iopub.execute_input":"2022-05-13T00:56:14.410098Z","iopub.status.idle":"2022-05-13T00:56:14.419379Z","shell.execute_reply.started":"2022-05-13T00:56:14.410049Z","shell.execute_reply":"2022-05-13T00:56:14.415555Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_wandb(file, run_path):\n    embs_cD = wandb.restore(file, run_path=run_path)\n\n    # # use the \"name\" attribute of the returned object if your framework expects a filename, e.g. as in Keras\n    embs_c = torch.load(embs_cD.name)\n    del embs_cD\n    return embs_c","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.420551Z","iopub.execute_input":"2022-05-13T00:56:14.420886Z","iopub.status.idle":"2022-05-13T00:56:14.435968Z","shell.execute_reply.started":"2022-05-13T00:56:14.420852Z","shell.execute_reply":"2022-05-13T00:56:14.435215Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# No need to dump shallow embdeing use as model layer\n# look at how long sequnece transaction sequence are... way long than RNN(Only Attention) \n# not possible else RNN+Attention ","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.455Z","iopub.execute_input":"2022-05-13T00:56:14.455253Z","iopub.status.idle":"2022-05-13T00:56:14.461781Z","shell.execute_reply.started":"2022-05-13T00:56:14.455222Z","shell.execute_reply":"2022-05-13T00:56:14.461135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Transformers Baby for Seq with Condition","metadata":{}},{"cell_type":"markdown","source":"# New Architect\nInput -  Articles[:10] Warm start encoder (10)\nNew Signal - Customer Emb Vect + Ecoded --> Encoded \nPass New Signal to Decoder Block\nOutputs -  Articles[10:]\n","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Custom Transformer to make context aware","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## End-To-End Transformer","metadata":{}},{"cell_type":"code","source":"###hollow","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.463058Z","iopub.execute_input":"2022-05-13T00:56:14.463564Z","iopub.status.idle":"2022-05-13T00:56:14.475161Z","shell.execute_reply.started":"2022-05-13T00:56:14.463526Z","shell.execute_reply":"2022-05-13T00:56:14.474493Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare Seq Data","metadata":{}},{"cell_type":"code","source":"import pickle\npkl_loc = '.'\nwith open(f'{pkl_loc}/le_a.pkl', 'rb') as fp:\n    le_a = pickle.load(fp)\n\nwith open(f'{pkl_loc}/le_c.pkl', 'rb') as fp:\n    le_c = pickle.load(fp)\nle_a.classes_ = le_a.classes_[:-2]\n\n# le_ad = wandb.restore(\"le_a.pkl\", \"mayankk-om-dev/HnmRGCNv2/34r1klbo\")\n# le_a = pickle.load(le_ad.name)\n\n# le_cd = wandb.restore(\"le_c.pkl\", \"mayankk-om-dev/HnmRGCNv2/34r1klbo\")\n# le_c = pickle.load(le_cd.name)\n\n# le_a.classes_ = le_a.classes_[:-2]","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.476213Z","iopub.execute_input":"2022-05-13T00:56:14.47655Z","iopub.status.idle":"2022-05-13T00:56:14.527287Z","shell.execute_reply.started":"2022-05-13T00:56:14.476518Z","shell.execute_reply":"2022-05-13T00:56:14.526534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import gc\n# gc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.528447Z","iopub.execute_input":"2022-05-13T00:56:14.529103Z","iopub.status.idle":"2022-05-13T00:56:14.5516Z","shell.execute_reply.started":"2022-05-13T00:56:14.52905Z","shell.execute_reply":"2022-05-13T00:56:14.550124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# transF = \"../input/h-and-m-personalized-fashion-recommendations/transactions_train.csv\"\n# df_trans = cudf.read_csv(transF)\n# df_trans[\"customer_id\"] = df_trans['customer_id'].str[-16:].str.hex_to_int().astype('int64')\n# df_trans['article_id'] = df_trans['article_id'].astype('int32')\n# df_trans['customer_id'] = le_c.transform(df_trans['customer_id'])\n# df_trans['article_id'] = le_a.transform(df_trans['article_id'])\n# df_trans = df_trans.to_pandas()\n# grp = df_trans.sort_values(['customer_id','t_dat']).groupby('customer_id')\n# gdd =grp['article_id'].apply(list)\n# gdd = gdd.reset_index()","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.555384Z","iopub.execute_input":"2022-05-13T00:56:14.55579Z","iopub.status.idle":"2022-05-13T00:56:14.564087Z","shell.execute_reply.started":"2022-05-13T00:56:14.555741Z","shell.execute_reply":"2022-05-13T00:56:14.563384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# gdd.head()\n# # transform ","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.566608Z","iopub.execute_input":"2022-05-13T00:56:14.571647Z","iopub.status.idle":"2022-05-13T00:56:14.582458Z","shell.execute_reply.started":"2022-05-13T00:56:14.571608Z","shell.execute_reply":"2022-05-13T00:56:14.581618Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# d = {'customer_id': [], 'article_src':[], 'article_dst':[]}\n\n# def split_row(row):\n# #     print(type(row['article_dst']))\n#     len_ = len(row['article_dst'])\n#     cust = row['customer_id']\n#     global d\n#     if len_<11:\n#         return row['article_dst']\n#     else:\n#         ret = row['article_dst'][:10] # max_seqsize =10\n#         extras = row['article_dst'][10:]\n#     while len(extras)>=15:\n#         # add to other df\n#         add = extras[:15]\n#         d['customer_id'].append(cust)\n#         d['article_src'].append(add[:5])\n#         d['article_dst'].append(add[5:])\n#         extras = extras[15:]\n    \n#     # add extras if extras and greater than 5\n# #     if extras and len(extras)>5:\n# #         d['customer_id'].append(cust)\n# #         d['article_src'].append(extras[:5])\n# #         d['article_dst'].append(extras[5:])\n\n#     return ret","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.586335Z","iopub.execute_input":"2022-05-13T00:56:14.586779Z","iopub.status.idle":"2022-05-13T00:56:14.596425Z","shell.execute_reply.started":"2022-05-13T00:56:14.586736Z","shell.execute_reply":"2022-05-13T00:56:14.59565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # remove less than 10 5enc 5 decode for train\n# gdd = gdd[gdd['article_id'].apply(lambda x: len(x))>15]\n# # create 5 -50+ \n# gdd['article_src'], gdd['article_dst'] = gdd['article_id'].apply(lambda x: x[:5]), gdd['article_id'].apply(lambda x: x[5:])\n# del gdd['article_id']\n# # split and create record on more than \n# gdd['article_dst'] = gdd.apply(split_row, axis=1)\n# df2 = pd.DataFrame.from_dict(d)\n# # df2 = df2[df2['article_dst'].apply(lambda x: len(x))>5]\n# fdf = pd.concat([df2, gdd])\n# del df2\n# del gdd","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.60199Z","iopub.execute_input":"2022-05-13T00:56:14.603101Z","iopub.status.idle":"2022-05-13T00:56:14.613467Z","shell.execute_reply.started":"2022-05-13T00:56:14.603051Z","shell.execute_reply":"2022-05-13T00:56:14.61261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# fdf['customer_id'] = fdf['customer_id'].astype(int)","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.615027Z","iopub.execute_input":"2022-05-13T00:56:14.615444Z","iopub.status.idle":"2022-05-13T00:56:14.622193Z","shell.execute_reply.started":"2022-05-13T00:56:14.61541Z","shell.execute_reply":"2022-05-13T00:56:14.621515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# fdf.dtypes","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.623506Z","iopub.execute_input":"2022-05-13T00:56:14.62493Z","iopub.status.idle":"2022-05-13T00:56:14.632539Z","shell.execute_reply.started":"2022-05-13T00:56:14.624891Z","shell.execute_reply":"2022-05-13T00:56:14.631361Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# fdf.to_csv(\"fdf.csv\", index=False)\n# wandb.save(\"fdf.csv\")\n# fdf.head(1000).to_csv(\"fdf2.csv\", index=False)\n# wandb.save(\"fdf2.csv\")","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.639243Z","iopub.execute_input":"2022-05-13T00:56:14.63971Z","iopub.status.idle":"2022-05-13T00:56:14.643973Z","shell.execute_reply.started":"2022-05-13T00:56:14.639677Z","shell.execute_reply":"2022-05-13T00:56:14.643311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fdf = wandb.restore(\"fdf.csv\", run_path=\"mayankk-om-dev/HnmRGCNv2/23rwj7i0\")","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:14.645271Z","iopub.execute_input":"2022-05-13T00:56:14.6458Z","iopub.status.idle":"2022-05-13T00:56:19.487503Z","shell.execute_reply.started":"2022-05-13T00:56:14.645761Z","shell.execute_reply":"2022-05-13T00:56:19.486605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(fdf.name)\nprint(df.shape)","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:19.488873Z","iopub.execute_input":"2022-05-13T00:56:19.489133Z","iopub.status.idle":"2022-05-13T00:56:22.616431Z","shell.execute_reply.started":"2022-05-13T00:56:19.489097Z","shell.execute_reply":"2022-05-13T00:56:22.605129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.to_csv(\"fdf.csv\", index=False)\nwandb.save(\"fdf.csv\")","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:22.622027Z","iopub.execute_input":"2022-05-13T00:56:22.628248Z","iopub.status.idle":"2022-05-13T00:56:30.573813Z","shell.execute_reply.started":"2022-05-13T00:56:22.628195Z","shell.execute_reply":"2022-05-13T00:56:30.573097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del df","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:30.574925Z","iopub.execute_input":"2022-05-13T00:56:30.575305Z","iopub.status.idle":"2022-05-13T00:56:30.719285Z","shell.execute_reply.started":"2022-05-13T00:56:30.575261Z","shell.execute_reply":"2022-05-13T00:56:30.71838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n# traindf, valdf = train_test_split(gd, train_size=0.8, random_state=29)\nfrom torchtext.legacy.data import Field, TabularDataset, BucketIterator\n\ncustomer = Field(sequential=False, use_vocab=False)\narticle_src = Field(sequential=True, use_vocab=False, tokenize=lambda x: eval(x))\narticle_dst = Field(sequential=True, use_vocab=False, pad_token=105540, fix_length=10, tokenize=lambda x: eval(x))\n\nfields = {'customer_id': ('c', customer), \n          'article_src':('asrc', article_src),\n          'article_dst':('atrg', article_dst)}\n\n\n","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:30.720772Z","iopub.execute_input":"2022-05-13T00:56:30.721377Z","iopub.status.idle":"2022-05-13T00:56:30.829061Z","shell.execute_reply.started":"2022-05-13T00:56:30.721339Z","shell.execute_reply":"2022-05-13T00:56:30.828149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# fdf.head()","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:30.833719Z","iopub.execute_input":"2022-05-13T00:56:30.834019Z","iopub.status.idle":"2022-05-13T00:56:30.842199Z","shell.execute_reply.started":"2022-05-13T00:56:30.833978Z","shell.execute_reply":"2022-05-13T00:56:30.841397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = TabularDataset(path='fdf.csv', format='csv', fields=fields)","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:56:30.844003Z","iopub.execute_input":"2022-05-13T00:56:30.846335Z","iopub.status.idle":"2022-05-13T00:57:35.493142Z","shell.execute_reply.started":"2022-05-13T00:56:30.846295Z","shell.execute_reply":"2022-05-13T00:57:35.492328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train[0].__dict__.keys()","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:57:35.494476Z","iopub.execute_input":"2022-05-13T00:57:35.494726Z","iopub.status.idle":"2022-05-13T00:57:35.501725Z","shell.execute_reply.started":"2022-05-13T00:57:35.494692Z","shell.execute_reply":"2022-05-13T00:57:35.501041Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train[-200].__dict__.values()#105540","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:57:35.503262Z","iopub.execute_input":"2022-05-13T00:57:35.503963Z","iopub.status.idle":"2022-05-13T00:57:35.519712Z","shell.execute_reply.started":"2022-05-13T00:57:35.503925Z","shell.execute_reply":"2022-05-13T00:57:35.51874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for batch in train_iterator:\n#     customers = batch.c\n#     customers = customers.view(1, customers.shape[0])\n#     print(customers.shape)\n#     break","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:57:35.521579Z","iopub.execute_input":"2022-05-13T00:57:35.521895Z","iopub.status.idle":"2022-05-13T00:57:35.528397Z","shell.execute_reply.started":"2022-05-13T00:57:35.521844Z","shell.execute_reply":"2022-05-13T00:57:35.527315Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train Model","metadata":{}},{"cell_type":"code","source":"from torch.nn.functional import linear, softmax, dropout\nimport math\nfrom torch.overrides import has_torch_function\nTensor = torch.Tensor\nfrom typing import Callable, List, Optional, Tuple, Union\ndef _in_projection_packed(\n    q: Tensor,\n    k: Tensor,\n    v: Tensor,\n    w: Tensor,\n    b: Optional[Tensor] = None,\n) -> List[Tensor]:\n    r\"\"\"\n    Performs the in-projection step of the attention operation, using packed weights.\n    Output is a triple containing projection tensors for query, key and value.\n    Args:\n        q, k, v: query, key and value tensors to be projected. For self-attention,\n            these are typically the same tensor; for encoder-decoder attention,\n            k and v are typically the same tensor. (We take advantage of these\n            identities for performance if they are present.) Regardless, q, k and v\n            must share a common embedding dimension; otherwise their shapes may vary.\n        w: projection weights for q, k and v, packed into a single tensor. Weights\n            are packed along dimension 0, in q, k, v order.\n        b: optional projection biases for q, k and v, packed into a single tensor\n            in q, k, v order.\n    Shape:\n        Inputs:\n        - q: :math:`(..., E)` where E is the embedding dimension\n        - k: :math:`(..., E)` where E is the embedding dimension\n        - v: :math:`(..., E)` where E is the embedding dimension\n        - w: :math:`(E * 3, E)` where E is the embedding dimension\n        - b: :math:`E * 3` where E is the embedding dimension\n        Output:\n        - in output list :math:`[q', k', v']`, each output tensor will have the\n            same shape as the corresponding input tensor.\n    \"\"\"\n    E = q.size(-1)\n    if k is v:\n        if q is k:\n            # self-attention\n            return linear(q, w, b).chunk(3, dim=-1)\n        else:\n            # encoder-decoder attention\n            w_q, w_kv = w.split([E, E * 2])\n            if b is None:\n                b_q = b_kv = None\n            else:\n                b_q, b_kv = b.split([E, E * 2])\n            return (linear(q, w_q, b_q),) + linear(k, w_kv, b_kv).chunk(2, dim=-1)\n    else:\n        w_q, w_k, w_v = w.chunk(3)\n        if b is None:\n            b_q = b_k = b_v = None\n        else:\n            b_q, b_k, b_v = b.chunk(3)\n        return linear(q, w_q, b_q), linear(k, w_k, b_k), linear(v, w_v, b_v)\n\ndef _scaled_dot_product_attention(\n    q: Tensor,\n    k: Tensor,\n    v: Tensor,\n    attn_mask: Optional[Tensor] = None,\n    dropout_p: float = 0.0,\n) -> Tuple[Tensor, Tensor]:\n    r\"\"\"\n    Computes scaled dot product attention on query, key and value tensors, using\n    an optional attention mask if passed, and applying dropout if a probability\n    greater than 0.0 is specified.\n    Returns a tensor pair containing attended values and attention weights.\n    Args:\n        q, k, v: query, key and value tensors. See Shape section for shape details.\n        attn_mask: optional tensor containing mask values to be added to calculated\n            attention. May be 2D or 3D; see Shape section for details.\n        dropout_p: dropout probability. If greater than 0.0, dropout is applied.\n    Shape:\n        - q: :math:`(B, Nt, E)` where B is batch size, Nt is the target sequence length,\n            and E is embedding dimension.\n        - key: :math:`(B, Ns, E)` where B is batch size, Ns is the source sequence length,\n            and E is embedding dimension.\n        - value: :math:`(B, Ns, E)` where B is batch size, Ns is the source sequence length,\n            and E is embedding dimension.\n        - attn_mask: either a 3D tensor of shape :math:`(B, Nt, Ns)` or a 2D tensor of\n            shape :math:`(Nt, Ns)`.\n        - Output: attention values have shape :math:`(B, Nt, E)`; attention weights\n            have shape :math:`(B, Nt, Ns)`\n    \"\"\"\n    B, Nt, E = q.shape\n    q = q / math.sqrt(E)\n    # (B, Nt, E) x (B, E, Ns) -> (B, Nt, Ns)\n    if attn_mask is not None:\n        attn = torch.baddbmm(attn_mask, q, k.transpose(-2, -1))\n    else:\n        attn = torch.bmm(q, k.transpose(-2, -1))\n#     print(\"Before Sftmx\", attn)\n    attn = softmax(attn, dim=-1)\n#     print(\"After Sftmx\", attn)\n    if dropout_p > 0.0:\n        attn = dropout(attn, p=dropout_p)\n    # (B, Nt, Ns) x (B, Ns, E) -> (B, Nt, E)\n    output = torch.bmm(attn, v)\n    return output, attn\n\ndef _in_projection(\n    q: Tensor,\n    k: Tensor,\n    v: Tensor,\n    w_q: Tensor,\n    w_k: Tensor,\n    w_v: Tensor,\n    b_q: Optional[Tensor] = None,\n    b_k: Optional[Tensor] = None,\n    b_v: Optional[Tensor] = None,\n) -> Tuple[Tensor, Tensor, Tensor]:\n    r\"\"\"\n    Performs the in-projection step of the attention operation. This is simply\n    a triple of linear projections, with shape constraints on the weights which\n    ensure embedding dimension uniformity in the projected outputs.\n    Output is a triple containing projection tensors for query, key and value.\n    Args:\n        q, k, v: query, key and value tensors to be projected.\n        w_q, w_k, w_v: weights for q, k and v, respectively.\n        b_q, b_k, b_v: optional biases for q, k and v, respectively.\n    Shape:\n        Inputs:\n        - q: :math:`(Qdims..., Eq)` where Eq is the query embedding dimension and Qdims are any\n            number of leading dimensions.\n        - k: :math:`(Kdims..., Ek)` where Ek is the key embedding dimension and Kdims are any\n            number of leading dimensions.\n        - v: :math:`(Vdims..., Ev)` where Ev is the value embedding dimension and Vdims are any\n            number of leading dimensions.\n        - w_q: :math:`(Eq, Eq)`\n        - w_k: :math:`(Eq, Ek)`\n        - w_v: :math:`(Eq, Ev)`\n        - b_q: :math:`(Eq)`\n        - b_k: :math:`(Eq)`\n        - b_v: :math:`(Eq)`\n        Output: in output triple :math:`(q', k', v')`,\n         - q': :math:`[Qdims..., Eq]`\n         - k': :math:`[Kdims..., Eq]`\n         - v': :math:`[Vdims..., Eq]`\n    \"\"\"\n    Eq, Ek, Ev = q.size(-1), k.size(-1), v.size(-1)\n    assert w_q.shape == (Eq, Eq), f\"expecting query weights shape of {(Eq, Eq)}, but got {w_q.shape}\"\n    assert w_k.shape == (Eq, Ek), f\"expecting key weights shape of {(Eq, Ek)}, but got {w_k.shape}\"\n    assert w_v.shape == (Eq, Ev), f\"expecting value weights shape of {(Eq, Ev)}, but got {w_v.shape}\"\n    assert b_q is None or b_q.shape == (Eq,), f\"expecting query bias shape of {(Eq,)}, but got {b_q.shape}\"\n    assert b_k is None or b_k.shape == (Eq,), f\"expecting key bias shape of {(Eq,)}, but got {b_k.shape}\"\n    assert b_v is None or b_v.shape == (Eq,), f\"expecting value bias shape of {(Eq,)}, but got {b_v.shape}\"\n    return linear(q, w_q, b_q), linear(k, w_k, b_k), linear(v, w_v, b_v)\n\ndef _mha_shape_check(query: Tensor, key: Tensor, value: Tensor,\n                     key_padding_mask: Optional[Tensor], attn_mask: Optional[Tensor], num_heads: int):\n    # Verifies the expected shape for `query, `key`, `value`, `key_padding_mask` and `attn_mask`\n    # and returns if the input is batched or not.\n    # Raises an error if `query` is not 2-D (unbatched) or 3-D (batched) tensor.\n\n    # Shape check.\n    if query.dim() == 3:\n        # Batched Inputs\n        is_batched = True\n        assert key.dim() == 3 and value.dim() == 3, \\\n            (\"For batched (3-D) `query`, expected `key` and `value` to be 3-D\"\n             f\" but found {key.dim()}-D and {value.dim()}-D tensors respectively\")\n        if key_padding_mask is not None:\n            assert key_padding_mask.dim() == 2, \\\n                (\"For batched (3-D) `query`, expected `key_padding_mask` to be `None` or 2-D\"\n                 f\" but found {key_padding_mask.dim()}-D tensor instead\")\n        if attn_mask is not None:\n            assert attn_mask.dim() in (2, 3), \\\n                (\"For batched (3-D) `query`, expected `attn_mask` to be `None`, 2-D or 3-D\"\n                 f\" but found {attn_mask.dim()}-D tensor instead\")\n    elif query.dim() == 2:\n        # Unbatched Inputs\n        is_batched = False\n        assert key.dim() == 2 and value.dim() == 2, \\\n            (\"For unbatched (2-D) `query`, expected `key` and `value` to be 2-D\"\n             f\" but found {key.dim()}-D and {value.dim()}-D tensors respectively\")\n\n        if key_padding_mask is not None:\n            assert key_padding_mask.dim() == 1, \\\n                (\"For unbatched (2-D) `query`, expected `key_padding_mask` to be `None` or 1-D\"\n                 f\" but found {key_padding_mask.dim()}-D tensor instead\")\n\n        if attn_mask is not None:\n            assert attn_mask.dim() in (2, 3), \\\n                (\"For unbatched (2-D) `query`, expected `attn_mask` to be `None`, 2-D or 3-D\"\n                 f\" but found {attn_mask.dim()}-D tensor instead\")\n            if attn_mask.dim() == 3:\n                expected_shape = (num_heads, query.shape[0], key.shape[0])\n                assert attn_mask.shape == expected_shape, \\\n                    (f\"Expected `attn_mask` shape to be {expected_shape} but got {attn_mask.shape}\")\n    else:\n        raise AssertionError(\n            f\"query should be unbatched 2D or batched 3D tensor but received {query.dim()}-D query tensor\")\n\n    return is_batched","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:57:35.530533Z","iopub.execute_input":"2022-05-13T00:57:35.530909Z","iopub.status.idle":"2022-05-13T00:57:35.575511Z","shell.execute_reply.started":"2022-05-13T00:57:35.530874Z","shell.execute_reply":"2022-05-13T00:57:35.574701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RecSTransformer(nn.Module):\n    def __init__(\n        self,\n        embedding_size,\n        src_vocab_size,\n        trg_vocab_size,\n        src_pad_idx,\n        num_heads,\n        num_encoder_layers,\n        num_decoder_layers,\n        dim_feedforward,\n        dropout,\n        max_len_s,\n        max_len_t,\n        device,\n    ):\n        super(RecSTransformer, self).__init__()\n#         weight_a = load_wandb(\"articleGNNEmb.pt\", \"mayankk-om-dev/HnmRGCNv2/2oooot9x\")\n#         weight_a = torch.cat([weight_a, torch.Tensor(1, 256)]) # included pad as last token\n        # self.article_embedding = nn.Embedding.from_pretrained(weight_a)\n        self.src_word_embedding = nn.Embedding(src_vocab_size, embedding_size)\n#         self.src_word_embedding = nn.Embedding.from_pretrained(weight_a) # maximize information for article\n        self.src_position_embedding = nn.Embedding(max_len_s+1, embedding_size)\n        # self.trg_word_embedding = nn.Embedding(trg_vocab_size, embedding_size)\n        weight_c = load_wandb(\"customerGNNEmb.pt\", \"mayankk-om-dev/HnmRGCNv2/2oooot9x\")\n        # self.trg_word_embedding = nn.Embedding.from_pretrained(weight_c)\n#         self.trg_word_embedding = self.src_word_embedding # nn.Embedding.from_pretrained(weight_a)#  maximize information for predicting much longer article seq\n        self.trg_word_embedding = nn.Embedding(max_len_t, embedding_size)\n        self.context_embedding = nn.Embedding(weight_c.shape[0], embedding_size) # nn.Embedding.from_pretrained(weight_c)# maximize information for predicting much longer article seq\n        self.trg_position_embedding = nn.Embedding(max_len_t, embedding_size)\n\n        self.device = device\n        self.max_len_s = max_len_s\n        self.max_len_t = max_len_t\n        self.transformer = Transformer(\n            embedding_size,\n            num_heads,\n            num_encoder_layers,\n            num_decoder_layers,\n            dim_feedforward,\n            dropout,\n            device=device\n        )\n        \n        self.fc_out = nn.Linear(embedding_size, trg_vocab_size)\n        self.dropout = nn.Dropout(dropout)\n        self.src_pad_idx = src_pad_idx\n\n    def make_src_mask(self, src):\n        src_mask = src.transpose(0, 1) == self.src_pad_idx\n\n        # (N, src_len)\n        return src_mask.to(self.device)\n    \n    def forward(self, src, trg, cust):\n        src_seq_length, N = src.shape\n        trg_seq_length, N = trg.shape\n        custID, N = cust.shape\n#         print(f\"Src {src.shape}, Tar {trg.shape}, Cus {cust.shape}\")\n        embed_cust = self.context_embedding(cust)\n#         print(embed_cust.shape)\n        embed_cust = embed_cust.repeat(self.max_len_s,1, 1) #N, seq_len, embs\n#       \n        src_positions = (\n            torch.arange(0, src_seq_length)\n            .unsqueeze(1)\n            .expand(src_seq_length, N)\n            .to(self.device)\n        )\n\n        trg_positions = (\n            torch.arange(0, trg_seq_length)\n            .unsqueeze(1)\n            .expand(trg_seq_length, N)\n            .to(self.device)\n        )\n#         print(f\"SrcP {src_positions.shape}, TarP {trg_positions.shape}\")\n#         print(self.src_word_embedding(src).shape)\n        embed_src = self.dropout(\n            (self.src_word_embedding(src) + self.src_position_embedding(src_positions))\n        )\n        embed_trg = self.dropout(\n            (self.trg_word_embedding(trg) + self.trg_position_embedding(trg_positions))\n        )\n#         print(\"SrcPosition\", self.src_position_embedding(src_positions))\n#         src_padding_mask = self.make_src_mask(src)\n        tgt_key_padding_mask = self.make_src_mask(trg)\n#         print(f\"embed_src {embed_src}, \\n embed_trg {embed_trg}\")\n        trg_mask = self.transformer.generate_square_subsequent_mask(trg_seq_length).to(\n            self.device\n        )\n#         embed_cust = self.context_embedding(cust)\n#         print(embed_cust.shape)\n#         embed_cust = embed_cust.repeat(self.max_len_s,1, 1) #N, seq_len, embs\n#         print(f\"embed_cust shape {embed_cust.shape}\")\n        out = self.transformer(\n            embed_src,\n            embed_trg,\n            embed_cust,\n#             src_key_padding_mask=src_padding_mask,\n            tgt_mask=trg_mask,\n            tgt_key_padding_mask = tgt_key_padding_mask\n        )\n        out = self.fc_out(out)\n        return out","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:57:35.577149Z","iopub.execute_input":"2022-05-13T00:57:35.577733Z","iopub.status.idle":"2022-05-13T00:57:35.599156Z","shell.execute_reply.started":"2022-05-13T00:57:35.577551Z","shell.execute_reply":"2022-05-13T00:57:35.598473Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef multi_head_attention_forward(\n    query: Tensor,\n    key: Tensor,\n    value: Tensor,\n    embed_dim_to_check: int,\n    num_heads: int,\n    in_proj_weight: Optional[Tensor],\n    in_proj_bias: Optional[Tensor],\n    bias_k: Optional[Tensor],\n    bias_v: Optional[Tensor],\n    add_zero_attn: bool,\n    dropout_p: float,\n    out_proj_weight: Tensor,\n    out_proj_bias: Optional[Tensor],\n    training: bool = True,\n    key_padding_mask: Optional[Tensor] = None,\n    need_weights: bool = True,\n    attn_mask: Optional[Tensor] = None,\n    use_separate_proj_weight: bool = False,\n    q_proj_weight: Optional[Tensor] = None,\n    k_proj_weight: Optional[Tensor] = None,\n    v_proj_weight: Optional[Tensor] = None,\n    static_k: Optional[Tensor] = None,\n    static_v: Optional[Tensor] = None,\n    average_attn_weights: bool = True,\n) -> Tuple[Tensor, Optional[Tensor]]:\n    r\"\"\"\n    Args:\n        query, key, value: map a query and a set of key-value pairs to an output.\n            See \"Attention Is All You Need\" for more details.\n        embed_dim_to_check: total dimension of the model.\n        num_heads: parallel attention heads.\n        in_proj_weight, in_proj_bias: input projection weight and bias.\n        bias_k, bias_v: bias of the key and value sequences to be added at dim=0.\n        add_zero_attn: add a new batch of zeros to the key and\n                       value sequences at dim=1.\n        dropout_p: probability of an element to be zeroed.\n        out_proj_weight, out_proj_bias: the output projection weight and bias.\n        training: apply dropout if is ``True``.\n        key_padding_mask: if provided, specified padding elements in the key will\n            be ignored by the attention. This is an binary mask. When the value is True,\n            the corresponding value on the attention layer will be filled with -inf.\n        need_weights: output attn_output_weights.\n        attn_mask: 2D or 3D mask that prevents attention to certain positions. A 2D mask will be broadcasted for all\n            the batches while a 3D mask allows to specify a different mask for the entries of each batch.\n        use_separate_proj_weight: the function accept the proj. weights for query, key,\n            and value in different forms. If false, in_proj_weight will be used, which is\n            a combination of q_proj_weight, k_proj_weight, v_proj_weight.\n        q_proj_weight, k_proj_weight, v_proj_weight, in_proj_bias: input projection weight and bias.\n        static_k, static_v: static key and value used for attention operators.\n        average_attn_weights: If true, indicates that the returned ``attn_weights`` should be averaged across heads.\n            Otherwise, ``attn_weights`` are provided separately per head. Note that this flag only has an effect\n            when ``need_weights=True.``. Default: True\n    Shape:\n        Inputs:\n        - query: :math:`(L, E)` or :math:`(L, N, E)` where L is the target sequence length, N is the batch size, E is\n          the embedding dimension.\n        - key: :math:`(S, E)` or :math:`(S, N, E)`, where S is the source sequence length, N is the batch size, E is\n          the embedding dimension.\n        - value: :math:`(S, E)` or :math:`(S, N, E)` where S is the source sequence length, N is the batch size, E is\n          the embedding dimension.\n        - key_padding_mask: :math:`(S)` or :math:`(N, S)` where N is the batch size, S is the source sequence length.\n          If a ByteTensor is provided, the non-zero positions will be ignored while the zero positions\n          will be unchanged. If a BoolTensor is provided, the positions with the\n          value of ``True`` will be ignored while the position with the value of ``False`` will be unchanged.\n        - attn_mask: 2D mask :math:`(L, S)` where L is the target sequence length, S is the source sequence length.\n          3D mask :math:`(N*num_heads, L, S)` where N is the batch size, L is the target sequence length,\n          S is the source sequence length. attn_mask ensures that position i is allowed to attend the unmasked\n          positions. If a ByteTensor is provided, the non-zero positions are not allowed to attend\n          while the zero positions will be unchanged. If a BoolTensor is provided, positions with ``True``\n          are not allowed to attend while ``False`` values will be unchanged. If a FloatTensor\n          is provided, it will be added to the attention weight.\n        - static_k: :math:`(N*num_heads, S, E/num_heads)`, where S is the source sequence length,\n          N is the batch size, E is the embedding dimension. E/num_heads is the head dimension.\n        - static_v: :math:`(N*num_heads, S, E/num_heads)`, where S is the source sequence length,\n          N is the batch size, E is the embedding dimension. E/num_heads is the head dimension.\n        Outputs:\n        - attn_output: :math:`(L, E)` or :math:`(L, N, E)` where L is the target sequence length, N is the batch size,\n          E is the embedding dimension.\n        - attn_output_weights: Only returned when ``need_weights=True``. If ``average_attn_weights=True``, returns\n          attention weights averaged across heads of shape :math:`(L, S)` when input is unbatched or\n          :math:`(N, L, S)`, where :math:`N` is the batch size, :math:`L` is the target sequence length, and\n          :math:`S` is the source sequence length. If ``average_weights=False``, returns attention weights per\n          head of shape :math:`(num_heads, L, S)` when input is unbatched or :math:`(N, num_heads, L, S)`.\n    \"\"\"\n    tens_ops = (query, key, value, in_proj_weight, in_proj_bias, bias_k, bias_v, out_proj_weight, out_proj_bias)\n    if has_torch_function(tens_ops):\n        return handle_torch_function(\n            multi_head_attention_forward,\n            tens_ops,\n            query,\n            key,\n            value,\n            embed_dim_to_check,\n            num_heads,\n            in_proj_weight,\n            in_proj_bias,\n            bias_k,\n            bias_v,\n            add_zero_attn,\n            dropout_p,\n            out_proj_weight,\n            out_proj_bias,\n            training=training,\n            key_padding_mask=key_padding_mask,\n            need_weights=need_weights,\n            attn_mask=attn_mask,\n            use_separate_proj_weight=use_separate_proj_weight,\n            q_proj_weight=q_proj_weight,\n            k_proj_weight=k_proj_weight,\n            v_proj_weight=v_proj_weight,\n            static_k=static_k,\n            static_v=static_v,\n            average_attn_weights=average_attn_weights,\n        )\n\n    is_batched = _mha_shape_check(query, key, value, key_padding_mask, attn_mask, num_heads)\n\n    # For unbatched input, we unsqueeze at the expected batch-dim to pretend that the input\n    # is batched, run the computation and before returning squeeze the\n    # batch dimension so that the output doesn't carry this temporary batch dimension.\n    if not is_batched:\n        # unsqueeze if the input is unbatched\n        query = query.unsqueeze(1)\n        key = key.unsqueeze(1)\n        value = value.unsqueeze(1)\n        if key_padding_mask is not None:\n            key_padding_mask = key_padding_mask.unsqueeze(0)\n\n    # set up shape vars\n    tgt_len, bsz, embed_dim = query.shape\n    src_len, _, _ = key.shape\n    assert embed_dim == embed_dim_to_check, \\\n        f\"was expecting embedding dimension of {embed_dim_to_check}, but got {embed_dim}\"\n    if isinstance(embed_dim, torch.Tensor):\n        # embed_dim can be a tensor when JIT tracing\n        head_dim = embed_dim.div(num_heads, rounding_mode='trunc')\n    else:\n        head_dim = embed_dim // num_heads\n    assert head_dim * num_heads == embed_dim, f\"embed_dim {embed_dim} not divisible by num_heads {num_heads}\"\n    if use_separate_proj_weight:\n        # allow MHA to have different embedding dimensions when separate projection weights are used\n        assert key.shape[:2] == value.shape[:2], \\\n            f\"key's sequence and batch dims {key.shape[:2]} do not match value's {value.shape[:2]}\"\n    else:\n        assert key.shape == value.shape, f\"key shape {key.shape} does not match value shape {value.shape}\"\n\n    #\n    # compute in-projection\n    #\n    if not use_separate_proj_weight:\n        assert in_proj_weight is not None, \"use_separate_proj_weight is False but in_proj_weight is None\"\n        q, k, v = _in_projection_packed(query, key, value, in_proj_weight, in_proj_bias)\n    else:\n        assert q_proj_weight is not None, \"use_separate_proj_weight is True but q_proj_weight is None\"\n        assert k_proj_weight is not None, \"use_separate_proj_weight is True but k_proj_weight is None\"\n        assert v_proj_weight is not None, \"use_separate_proj_weight is True but v_proj_weight is None\"\n        if in_proj_bias is None:\n            b_q = b_k = b_v = None\n        else:\n            b_q, b_k, b_v = in_proj_bias.chunk(3)\n        q, k, v = _in_projection(query, key, value, q_proj_weight, k_proj_weight, v_proj_weight, b_q, b_k, b_v)\n\n    # prep attention mask\n    if attn_mask is not None:\n        if attn_mask.dtype == torch.uint8:\n            warnings.warn(\"Byte tensor for attn_mask in nn.MultiheadAttention is deprecated. Use bool tensor instead.\")\n            attn_mask = attn_mask.to(torch.bool)\n        else:\n            assert attn_mask.is_floating_point() or attn_mask.dtype == torch.bool, \\\n                f\"Only float, byte, and bool types are supported for attn_mask, not {attn_mask.dtype}\"\n        # ensure attn_mask's dim is 3\n        if attn_mask.dim() == 2:\n            correct_2d_size = (tgt_len, src_len)\n            if attn_mask.shape != correct_2d_size:\n                raise RuntimeError(f\"The shape of the 2D attn_mask is {attn_mask.shape}, but should be {correct_2d_size}.\")\n            attn_mask = attn_mask.unsqueeze(0)\n        elif attn_mask.dim() == 3:\n            correct_3d_size = (bsz * num_heads, tgt_len, src_len)\n            if attn_mask.shape != correct_3d_size:\n                raise RuntimeError(f\"The shape of the 3D attn_mask is {attn_mask.shape}, but should be {correct_3d_size}.\")\n        else:\n            raise RuntimeError(f\"attn_mask's dimension {attn_mask.dim()} is not supported\")\n\n    # prep key padding mask\n    if key_padding_mask is not None and key_padding_mask.dtype == torch.uint8:\n        warnings.warn(\"Byte tensor for key_padding_mask in nn.MultiheadAttention is deprecated. Use bool tensor instead.\")\n        key_padding_mask = key_padding_mask.to(torch.bool)\n\n    # add bias along batch dimension (currently second)\n    if bias_k is not None and bias_v is not None:\n        assert static_k is None, \"bias cannot be added to static key.\"\n        assert static_v is None, \"bias cannot be added to static value.\"\n        k = torch.cat([k, bias_k.repeat(1, bsz, 1)])\n        v = torch.cat([v, bias_v.repeat(1, bsz, 1)])\n        if attn_mask is not None:\n            attn_mask = pad(attn_mask, (0, 1))\n            print(\"Padded Attn_Mask\",  attn_mask)\n        if key_padding_mask is not None:\n            key_padding_mask = pad(key_padding_mask, (0, 1))\n    else:\n        assert bias_k is None\n        assert bias_v is None\n\n    #\n    # reshape q, k, v for multihead attention and make em batch first\n    #\n    q = q.contiguous().view(tgt_len, bsz * num_heads, head_dim).transpose(0, 1)\n    if static_k is None:\n        k = k.contiguous().view(k.shape[0], bsz * num_heads, head_dim).transpose(0, 1)\n    else:\n        # TODO finish disentangling control flow so we don't do in-projections when statics are passed\n        assert static_k.size(0) == bsz * num_heads, \\\n            f\"expecting static_k.size(0) of {bsz * num_heads}, but got {static_k.size(0)}\"\n        assert static_k.size(2) == head_dim, \\\n            f\"expecting static_k.size(2) of {head_dim}, but got {static_k.size(2)}\"\n        k = static_k\n    if static_v is None:\n        v = v.contiguous().view(v.shape[0], bsz * num_heads, head_dim).transpose(0, 1)\n    else:\n        # TODO finish disentangling control flow so we don't do in-projections when statics are passed\n        assert static_v.size(0) == bsz * num_heads, \\\n            f\"expecting static_v.size(0) of {bsz * num_heads}, but got {static_v.size(0)}\"\n        assert static_v.size(2) == head_dim, \\\n            f\"expecting static_v.size(2) of {head_dim}, but got {static_v.size(2)}\"\n        v = static_v\n\n    # add zero attention along batch dimension (now first)\n    if add_zero_attn:\n        zero_attn_shape = (bsz * num_heads, 1, head_dim)\n        k = torch.cat([k, torch.zeros(zero_attn_shape, dtype=k.dtype, device=k.device)], dim=1)\n        v = torch.cat([v, torch.zeros(zero_attn_shape, dtype=v.dtype, device=v.device)], dim=1)\n        if attn_mask is not None:\n            attn_mask = pad(attn_mask, (0, 1))\n        if key_padding_mask is not None:\n            key_padding_mask = pad(key_padding_mask, (0, 1))\n\n    # update source sequence length after adjustments\n    src_len = k.size(1)\n\n    # merge key padding and attention masks\n    if key_padding_mask is not None:\n        assert key_padding_mask.shape == (bsz, src_len), \\\n            f\"expecting key_padding_mask shape of {(bsz, src_len)}, but got {key_padding_mask.shape}\"\n        key_padding_mask = key_padding_mask.view(bsz, 1, 1, src_len).   \\\n            expand(-1, num_heads, -1, -1).reshape(bsz * num_heads, 1, src_len)\n        if attn_mask is None:\n            attn_mask = key_padding_mask\n        elif attn_mask.dtype == torch.bool:\n            attn_mask = attn_mask.logical_or(key_padding_mask)\n        else:\n            attn_mask = attn_mask.masked_fill(key_padding_mask, float(\"-1e20\"))\n    \n    # convert mask to float\n    if attn_mask is not None and attn_mask.dtype == torch.bool:\n        new_attn_mask = torch.zeros_like(attn_mask, dtype=q.dtype)\n        new_attn_mask.masked_fill_(attn_mask, float(\"-1e20\"))\n        attn_mask = new_attn_mask\n#     print(\"Attention Mask\", attn_mask )\n    # adjust dropout probability\n    if not training:\n        dropout_p = 0.0\n\n    #\n    # (deep breath) calculate attention and out projection\n    #\n    attn_output, attn_output_weights = _scaled_dot_product_attention(q, k, v, attn_mask, dropout_p)\n    attn_output = attn_output.transpose(0, 1).contiguous().view(tgt_len * bsz, embed_dim)\n    attn_output = linear(attn_output, out_proj_weight, out_proj_bias)\n    attn_output = attn_output.view(tgt_len, bsz, attn_output.size(1))\n\n    if need_weights:\n        # optionally average attention weights over heads\n        attn_output_weights = attn_output_weights.view(bsz, num_heads, tgt_len, src_len)\n        if average_attn_weights:\n            attn_output_weights = attn_output_weights.sum(dim=1) / num_heads\n\n        if not is_batched:\n            # squeeze the output if input was unbatched\n            attn_output = attn_output.squeeze(1)\n            attn_output_weights = attn_output_weights.squeeze(0)\n        return attn_output, attn_output_weights\n    else:\n        if not is_batched:\n            # squeeze the output if input was unbatched\n            attn_output = attn_output.squeeze(1)\n        return attn_output, None","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:57:35.602365Z","iopub.execute_input":"2022-05-13T00:57:35.602558Z","iopub.status.idle":"2022-05-13T00:57:35.650705Z","shell.execute_reply.started":"2022-05-13T00:57:35.602535Z","shell.execute_reply":"2022-05-13T00:57:35.649799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"linear","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:57:35.652277Z","iopub.execute_input":"2022-05-13T00:57:35.653098Z","iopub.status.idle":"2022-05-13T00:57:35.665858Z","shell.execute_reply.started":"2022-05-13T00:57:35.65304Z","shell.execute_reply":"2022-05-13T00:57:35.665124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import warnings\nfrom typing import Optional, Tuple\n\nimport torch\nfrom torch import Tensor\nfrom torch.nn.modules.linear import NonDynamicallyQuantizableLinear\nfrom torch.nn.init import constant_, xavier_normal_, xavier_uniform_\nfrom torch.nn.parameter import Parameter\nfrom torch.nn import Module\n\nclass MultiheadAttention(Module):\n    r\"\"\"Allows the model to jointly attend to information\n    from different representation subspaces as described in the paper:\n    `Attention Is All You Need <https://arxiv.org/abs/1706.03762>`_.\n\n    Multi-Head Attention is defined as:\n\n    .. math::\n        \\text{MultiHead}(Q, K, V) = \\text{Concat}(head_1,\\dots,head_h)W^O\n\n    where :math:`head_i = \\text{Attention}(QW_i^Q, KW_i^K, VW_i^V)`.\n\n    ``forward()`` will use a special optimized implementation if all of the following\n    conditions are met:\n\n    - self attention is being computed (i.e., ``query``, ``key``, and ``value`` are the same tensor. This\n      restriction will be loosened in the future.)\n    - Either autograd is disabled (using ``torch.inference_mode`` or ``torch.no_grad``) or no tensor argument ``requires_grad``\n    - training is disabled (using ``.eval()``)\n    - dropout is 0\n    - ``add_bias_kv`` is ``False``\n    - ``add_zero_attn`` is ``False``\n    - ``batch_first`` is ``True`` and the input is batched\n    - ``kdim`` and ``vdim`` are equal to ``embed_dim``\n    - at most one of ``key_padding_mask`` or ``attn_mask`` is passed\n    - if a `NestedTensor <https://pytorch.org/docs/stable/nested.html>`_ is passed, neither ``key_padding_mask``\n      nor ``attn_mask`` is passed\n\n    If the optimized implementation is in use, a\n    `NestedTensor <https://pytorch.org/docs/stable/nested.html>`_ can be passed for\n    ``query``/``key``/``value`` to represent padding more efficiently than using a\n    padding mask. In this case, a `NestedTensor <https://pytorch.org/docs/stable/nested.html>`_\n    will be returned, and an additional speedup proportional to the fraction of the input\n    that is padding can be expected.\n\n    Args:\n        embed_dim: Total dimension of the model.\n        num_heads: Number of parallel attention heads. Note that ``embed_dim`` will be split\n            across ``num_heads`` (i.e. each head will have dimension ``embed_dim // num_heads``).\n        dropout: Dropout probability on ``attn_output_weights``. Default: ``0.0`` (no dropout).\n        bias: If specified, adds bias to input / output projection layers. Default: ``True``.\n        add_bias_kv: If specified, adds bias to the key and value sequences at dim=0. Default: ``False``.\n        add_zero_attn: If specified, adds a new batch of zeros to the key and value sequences at dim=1.\n            Default: ``False``.\n        kdim: Total number of features for keys. Default: ``None`` (uses ``kdim=embed_dim``).\n        vdim: Total number of features for values. Default: ``None`` (uses ``vdim=embed_dim``).\n        batch_first: If ``True``, then the input and output tensors are provided\n            as (batch, seq, feature). Default: ``False`` (seq, batch, feature).\n\n    Examples::\n\n        >>> multihead_attn = nn.MultiheadAttention(embed_dim, num_heads)\n        >>> attn_output, attn_output_weights = multihead_attn(query, key, value)\n\n    \"\"\"\n    __constants__ = ['batch_first']\n    bias_k: Optional[torch.Tensor]\n    bias_v: Optional[torch.Tensor]\n\n    def __init__(self, embed_dim, num_heads, dropout=0., bias=True, add_bias_kv=False, add_zero_attn=False,\n                 kdim=None, vdim=None, batch_first=False, device=None, dtype=None) -> None:\n        factory_kwargs = {'device': device, 'dtype': dtype}\n        super(MultiheadAttention, self).__init__()\n        self.embed_dim = embed_dim\n        self.kdim = kdim if kdim is not None else embed_dim\n        self.vdim = vdim if vdim is not None else embed_dim\n        self._qkv_same_embed_dim = self.kdim == embed_dim and self.vdim == embed_dim\n\n        self.num_heads = num_heads\n        self.dropout = dropout\n        self.batch_first = batch_first\n        self.head_dim = embed_dim // num_heads\n        assert self.head_dim * num_heads == self.embed_dim, \"embed_dim must be divisible by num_heads\"\n\n        if self._qkv_same_embed_dim is False:\n            self.q_proj_weight = Parameter(torch.empty((embed_dim, embed_dim), **factory_kwargs))\n            self.k_proj_weight = Parameter(torch.empty((embed_dim, self.kdim), **factory_kwargs))\n            self.v_proj_weight = Parameter(torch.empty((embed_dim, self.vdim), **factory_kwargs))\n            self.register_parameter('in_proj_weight', None)\n        else:\n            self.in_proj_weight = Parameter(torch.empty((3 * embed_dim, embed_dim), **factory_kwargs))\n            self.register_parameter('q_proj_weight', None)\n            self.register_parameter('k_proj_weight', None)\n            self.register_parameter('v_proj_weight', None)\n\n        if bias:\n            self.in_proj_bias = Parameter(torch.empty(3 * embed_dim, **factory_kwargs))\n        else:\n            self.register_parameter('in_proj_bias', None)\n        self.out_proj = NonDynamicallyQuantizableLinear(embed_dim, embed_dim, bias=bias, **factory_kwargs)\n\n        if add_bias_kv:\n            self.bias_k = Parameter(torch.empty((1, 1, embed_dim), **factory_kwargs))\n            self.bias_v = Parameter(torch.empty((1, 1, embed_dim), **factory_kwargs))\n        else:\n            self.bias_k = self.bias_v = None\n\n        self.add_zero_attn = add_zero_attn\n\n        self._reset_parameters()\n\n    def _reset_parameters(self):\n        if self._qkv_same_embed_dim:\n            xavier_uniform_(self.in_proj_weight)\n        else:\n            xavier_uniform_(self.q_proj_weight)\n            xavier_uniform_(self.k_proj_weight)\n            xavier_uniform_(self.v_proj_weight)\n\n        if self.in_proj_bias is not None:\n            constant_(self.in_proj_bias, 0.)\n            constant_(self.out_proj.bias, 0.)\n        if self.bias_k is not None:\n            xavier_normal_(self.bias_k)\n        if self.bias_v is not None:\n            xavier_normal_(self.bias_v)\n\n    def __setstate__(self, state):\n        # Support loading old MultiheadAttention checkpoints generated by v1.1.0\n        if '_qkv_same_embed_dim' not in state:\n            state['_qkv_same_embed_dim'] = True\n\n        super(MultiheadAttention, self).__setstate__(state)\n#     (x, tgt_mask, tgt_key_padding_mask)\n    def forward(self, query: Tensor, key: Tensor, value: Tensor, key_padding_mask: Optional[Tensor] = None,\n                need_weights: bool = True, attn_mask: Optional[Tensor] = None,\n                average_attn_weights: bool = True) -> Tuple[Tensor, Optional[Tensor]]:\n        r\"\"\"\n    Args:\n        query: Query embeddings of shape :math:`(L, E_q)` for unbatched input, :math:`(L, N, E_q)` when ``batch_first=False``\n            or :math:`(N, L, E_q)` when ``batch_first=True``, where :math:`L` is the target sequence length,\n            :math:`N` is the batch size, and :math:`E_q` is the query embedding dimension ``embed_dim``.\n            Queries are compared against key-value pairs to produce the output.\n            See \"Attention Is All You Need\" for more details.\n        key: Key embeddings of shape :math:`(S, E_k)` for unbatched input, :math:`(S, N, E_k)` when ``batch_first=False``\n            or :math:`(N, S, E_k)` when ``batch_first=True``, where :math:`S` is the source sequence length,\n            :math:`N` is the batch size, and :math:`E_k` is the key embedding dimension ``kdim``.\n            See \"Attention Is All You Need\" for more details.\n        value: Value embeddings of shape :math:`(S, E_v)` for unbatched input, :math:`(S, N, E_v)` when\n            ``batch_first=False`` or :math:`(N, S, E_v)` when ``batch_first=True``, where :math:`S` is the source\n            sequence length, :math:`N` is the batch size, and :math:`E_v` is the value embedding dimension ``vdim``.\n            See \"Attention Is All You Need\" for more details.\n        key_padding_mask: If specified, a mask of shape :math:`(N, S)` indicating which elements within ``key``\n            to ignore for the purpose of attention (i.e. treat as \"padding\"). For unbatched `query`, shape should be :math:`(S)`.\n            Binary and byte masks are supported.\n            For a binary mask, a ``True`` value indicates that the corresponding ``key`` value will be ignored for\n            the purpose of attention. For a byte mask, a non-zero value indicates that the corresponding ``key``\n            value will be ignored.\n        need_weights: If specified, returns ``attn_output_weights`` in addition to ``attn_outputs``.\n            Default: ``True``.\n        attn_mask: If specified, a 2D or 3D mask preventing attention to certain positions. Must be of shape\n            :math:`(L, S)` or :math:`(N\\cdot\\text{num\\_heads}, L, S)`, where :math:`N` is the batch size,\n            :math:`L` is the target sequence length, and :math:`S` is the source sequence length. A 2D mask will be\n            broadcasted across the batch while a 3D mask allows for a different mask for each entry in the batch.\n            Binary, byte, and float masks are supported. For a binary mask, a ``True`` value indicates that the\n            corresponding position is not allowed to attend. For a byte mask, a non-zero value indicates that the\n            corresponding position is not allowed to attend. For a float mask, the mask values will be added to\n            the attention weight.\n        average_attn_weights: If true, indicates that the returned ``attn_weights`` should be averaged across\n            heads. Otherwise, ``attn_weights`` are provided separately per head. Note that this flag only has an\n            effect when ``need_weights=True``. Default: ``True`` (i.e. average weights across heads)\n\n    Outputs:\n        - **attn_output** - Attention outputs of shape :math:`(L, E)` when input is unbatched,\n          :math:`(L, N, E)` when ``batch_first=False`` or :math:`(N, L, E)` when ``batch_first=True``,\n          where :math:`L` is the target sequence length, :math:`N` is the batch size, and :math:`E` is the\n          embedding dimension ``embed_dim``.\n        - **attn_output_weights** - Only returned when ``need_weights=True``. If ``average_attn_weights=True``,\n          returns attention weights averaged across heads of shape :math:`(L, S)` when input is unbatched or\n          :math:`(N, L, S)`, where :math:`N` is the batch size, :math:`L` is the target sequence length, and\n          :math:`S` is the source sequence length. If ``average_weights=False``, returns attention weights per\n          head of shape :math:`(\\text{num\\_heads}, L, S)` when input is unbatched or :math:`(N, \\text{num\\_heads}, L, S)`.\n\n        .. note::\n            `batch_first` argument is ignored for unbatched inputs.\n        \"\"\"\n        is_batched = query.dim() == 3\n        why_not_fast_path = ''\n        if not is_batched:\n            why_not_fast_path = f\"input not batched; expected query.dim() of 3 but got {query.dim()}\"\n        elif query is not key or key is not value:\n            # When lifting this restriction, don't forget to either\n            # enforce that the dtypes all match or test cases where\n            # they don't!\n            why_not_fast_path = \"non-self attention was used (query, key, and value are not the same Tensor)\"\n        elif query.dtype != self.in_proj_bias.dtype:\n            why_not_fast_path = f\"dtypes of query ({query.dtype}) and self.in_proj_bias ({self.in_proj_bias.dtype}) don't match\"\n        elif query.dtype != self.in_proj_weight.dtype:\n            # this case will fail anyway, but at least they'll get a useful error message.\n            why_not_fast_path = f\"dtypes of query ({query.dtype}) and self.in_proj_weight ({self.in_proj_weight.dtype}) don't match\"\n        elif self.training:\n            why_not_fast_path = \"training is enabled\"\n        elif not self.batch_first:\n            why_not_fast_path = \"batch_first was not True\"\n        elif self.bias_k is not None:\n            why_not_fast_path = \"self.bias_k was not None\"\n        elif self.bias_v is not None:\n            why_not_fast_path = \"self.bias_v was not None\"\n        elif self.dropout:\n            why_not_fast_path = f\"dropout was {self.dropout}, required zero\"\n        elif self.add_zero_attn:\n            why_not_fast_path = \"add_zero_attn was enabled\"\n        elif not self._qkv_same_embed_dim:\n            why_not_fast_path = \"_qkv_same_embed_dim was not True\"\n        elif query.nest.is_nested and (key_padding_mask is not None or attn_mask is not None):\n            why_not_fast_path = \"key_padding_mask and attn_mask are not supported with NestedTensor input\"\n        elif not query.nest.is_nested and key_padding_mask is not None and attn_mask is not None:\n            why_not_fast_path = \"key_padding_mask and attn_mask were both supplied\"\n\n        if not why_not_fast_path:\n            tensor_args = (\n                query,\n                key,\n                value,\n                self.in_proj_weight,\n                self.in_proj_bias,\n                self.out_proj.weight,\n                self.out_proj.bias,\n            )\n            # We have to use list comprehensions below because TorchScript does not support\n            # generator expressions.\n            if torch.overrides.has_torch_function(tensor_args):\n                why_not_fast_path = \"some Tensor argument has_torch_function\"\n            elif not all([(x.is_cuda or 'cpu' in str(x.device)) for x in tensor_args]):\n                why_not_fast_path = \"some Tensor argument is neither CUDA nor CPU\"\n            elif torch.is_grad_enabled() and any([x.requires_grad for x in tensor_args]):\n                why_not_fast_path = (\"grad is enabled and at least one of query or the \"\n                                     \"input/output projection weights or biases requires_grad\")\n            if not why_not_fast_path:\n                return torch._native_multi_head_attention(\n                    query,\n                    key,\n                    value,\n                    self.embed_dim,\n                    self.num_heads,\n                    self.in_proj_weight,\n                    self.in_proj_bias,\n                    self.out_proj.weight,\n                    self.out_proj.bias,\n                    key_padding_mask if key_padding_mask is not None else attn_mask,\n                    need_weights,\n                    average_attn_weights)\n#         any_nested = query.is_nested or key.is_nested or value.is_nested\n#         assert not any_nested, (\"MultiheadAttention does not support NestedTensor outside of its fast path. \" +\n#                                 f\"The fast path was not hit because {why_not_fast_path}\")\n\n        if self.batch_first and is_batched:\n            # make sure that the transpose op does not affect the \"is\" property\n            if key is value:\n                if query is key:\n                    query = key = value = query.transpose(1, 0)\n                else:\n                    query, key = [x.transpose(1, 0) for x in (query, key)]\n                    value = key\n            else:\n                query, key, value = [x.transpose(1, 0) for x in (query, key, value)]\n\n        if not self._qkv_same_embed_dim:\n            attn_output, attn_output_weights = multi_head_attention_forward(\n                query, key, value, self.embed_dim, self.num_heads,\n                self.in_proj_weight, self.in_proj_bias,\n                self.bias_k, self.bias_v, self.add_zero_attn,\n                self.dropout, self.out_proj.weight, self.out_proj.bias,\n                training=self.training,\n                key_padding_mask=key_padding_mask, need_weights=need_weights,\n                attn_mask=attn_mask, use_separate_proj_weight=True,\n                q_proj_weight=self.q_proj_weight, k_proj_weight=self.k_proj_weight,\n                v_proj_weight=self.v_proj_weight, average_attn_weights=average_attn_weights)\n        else:\n            attn_output, attn_output_weights = multi_head_attention_forward(\n                query, key, value, self.embed_dim, self.num_heads,\n                self.in_proj_weight, self.in_proj_bias,\n                self.bias_k, self.bias_v, self.add_zero_attn,\n                self.dropout, self.out_proj.weight, self.out_proj.bias,\n                training=self.training,\n                key_padding_mask=key_padding_mask, need_weights=need_weights,\n                attn_mask=attn_mask, average_attn_weights=average_attn_weights)\n        if self.batch_first and is_batched:\n            return attn_output.transpose(1, 0), attn_output_weights\n        else:\n            return attn_output, attn_output_weights\n\n","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:57:35.667524Z","iopub.execute_input":"2022-05-13T00:57:35.667924Z","iopub.status.idle":"2022-05-13T00:57:35.722234Z","shell.execute_reply.started":"2022-05-13T00:57:35.667865Z","shell.execute_reply":"2022-05-13T00:57:35.721229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import copy\nfrom typing import Optional, Any, Union, Callable\n\nimport torch\nimport torch.nn as nn\nfrom torch import Tensor\nimport torch.nn.functional as F\nfrom torch.nn import Module, ModuleList, TransformerEncoderLayer, TransformerEncoder, TransformerDecoder\nfrom torch.nn.init import xavier_uniform_\nfrom torch.nn import Dropout, Linear, LayerNorm\n\nclass TransformerDecoderLayer(Module):\n    r\"\"\"\n    Args:\n        d_model: the number of expected features in the input (required).\n        nhead: the number of heads in the multiheadattention models (required).\n        dim_feedforward: the dimension of the feedforward network model (default=2048).\n        dropout: the dropout value (default=0.1).\n        activation: the activation function of the intermediate layer, can be a string\n            (\"relu\" or \"gelu\") or a unary callable. Default: relu\n        layer_norm_eps: the eps value in layer normalization components (default=1e-5).\n        batch_first: If ``True``, then the input and output tensors are provided\n            as (batch, seq, feature). Default: ``False`` (seq, batch, feature).\n        norm_first: if ``True``, layer norm is done prior to self attention, multihead\n            attention and feedforward operations, respectivaly. Otherwise it's done after.\n            Default: ``False`` (after).\n\n    Examples::\n        >>> decoder_layer = nn.TransformerDecoderLayer(d_model=512, nhead=8)\n        >>> memory = torch.rand(10, 32, 512)\n        >>> tgt = torch.rand(20, 32, 512)\n        >>> out = decoder_layer(tgt, memory)\n\n    Alternatively, when ``batch_first`` is ``True``:\n        >>> decoder_layer = nn.TransformerDecoderLayer(d_model=512, nhead=8, batch_first=True)\n        >>> memory = torch.rand(32, 10, 512)\n        >>> tgt = torch.rand(32, 20, 512)\n        >>> out = decoder_layer(tgt, memory)\n    \"\"\"\n    __constants__ = ['batch_first', 'norm_first']\n\n    def __init__(self, d_model: int, nhead: int, dim_feedforward: int = 2048, dropout: float = 0.1,\n                 activation: Union[str, Callable[[Tensor], Tensor]] = F.relu,\n                 layer_norm_eps: float = 1e-5, batch_first: bool = False, norm_first: bool = False,\n                 device=None, dtype=None) -> None:\n        factory_kwargs = {'device': device, 'dtype': dtype}\n        super(TransformerDecoderLayer, self).__init__()\n        self.self_attn = MultiheadAttention(d_model, nhead, dropout=dropout, batch_first=batch_first,\n                                            **factory_kwargs)\n        self.multihead_attn = MultiheadAttention(d_model, nhead, dropout=dropout, batch_first=batch_first,\n                                                 **factory_kwargs)\n        # Implementation of Feedforward model\n        self.linear1 = Linear(d_model, dim_feedforward, **factory_kwargs)\n        self.dropout = Dropout(dropout)\n        self.linear2 = Linear(dim_feedforward, d_model, **factory_kwargs)\n\n        self.norm_first = norm_first\n        self.norm1 = LayerNorm(d_model, eps=layer_norm_eps, **factory_kwargs)\n        self.norm2 = LayerNorm(d_model, eps=layer_norm_eps, **factory_kwargs)\n        self.norm3 = LayerNorm(d_model, eps=layer_norm_eps, **factory_kwargs)\n        self.dropout1 = Dropout(dropout)\n        self.dropout2 = Dropout(dropout)\n        self.dropout3 = Dropout(dropout)\n\n        # Legacy string support for activation function.\n        if isinstance(activation, str):\n            self.activation = _get_activation_fn(activation)\n        else:\n            self.activation = activation\n\n    def __setstate__(self, state):\n        if 'activation' not in state:\n            state['activation'] = F.relu\n        super(TransformerDecoderLayer, self).__setstate__(state)\n\n    def forward(self, tgt: Tensor, memory: Tensor, tgt_mask: Optional[Tensor] = None, memory_mask: Optional[Tensor] = None,\n                tgt_key_padding_mask: Optional[Tensor] = None, memory_key_padding_mask: Optional[Tensor] = None) -> Tensor:\n        r\"\"\"Pass the inputs (and mask) through the decoder layer.\n\n        Args:\n            tgt: the sequence to the decoder layer (required).\n            memory: the sequence from the last layer of the encoder (required).\n            tgt_mask: the mask for the tgt sequence (optional).\n            memory_mask: the mask for the memory sequence (optional).\n            tgt_key_padding_mask: the mask for the tgt keys per batch (optional).\n            memory_key_padding_mask: the mask for the memory keys per batch (optional).\n\n        Shape:\n            see the docs in Transformer class.\n        \"\"\"\n        # see Fig. 1 of https://arxiv.org/pdf/2002.04745v1.pdf\n\n        x = tgt\n        if self.norm_first:\n            x = x + self._sa_block(self.norm1(x), tgt_mask, tgt_key_padding_mask)\n            x = x + self._mha_block(self.norm2(x), memory, memory_mask, memory_key_padding_mask)\n            x = x + self._ff_block(self.norm3(x))\n        else:\n            x = self.norm1(x + self._sa_block(x, tgt_mask, tgt_key_padding_mask))\n#             print(\"selfAttention\", x)\n            x = self.norm2(x + self._mha_block(x, memory, memory_mask, memory_key_padding_mask))\n#             print(\"NormMHA\", x)\n            x = self.norm3(x + self._ff_block(x))\n#             print(\"Norm3\", x)\n#         print(x.shape)\n        return x\n\n\n    # self-attention block\n    def _sa_block(self, x: Tensor,\n                  attn_mask: Optional[Tensor], key_padding_mask: Optional[Tensor]) -> Tensor:\n        x = self.self_attn(x, x, x,\n                           attn_mask=attn_mask,\n                           key_padding_mask=key_padding_mask,\n                           need_weights=False)[0]\n        return self.dropout1(x)\n\n    # multihead attention block\n    def _mha_block(self, x: Tensor, mem: Tensor,\n                   attn_mask: Optional[Tensor], key_padding_mask: Optional[Tensor]) -> Tensor:\n        x = self.multihead_attn(x, mem, mem,\n                                attn_mask=attn_mask,\n                                key_padding_mask=key_padding_mask,\n                                need_weights=False)[0]\n        return self.dropout2(x)\n\n    # feed forward block\n    def _ff_block(self, x: Tensor) -> Tensor:\n        x = self.linear2(self.dropout(self.activation(self.linear1(x))))\n        return self.dropout3(x)\n\n\n\ndef _get_clones(module, N):\n    return ModuleList([copy.deepcopy(module) for i in range(N)])\n\n\ndef _get_activation_fn(activation):\n    if activation == \"relu\":\n        return F.relu\n    elif activation == \"gelu\":\n        return F.gelu\n\n    raise RuntimeError(\"activation should be relu/gelu, not {}\".format(activation))","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:57:35.724151Z","iopub.execute_input":"2022-05-13T00:57:35.724664Z","iopub.status.idle":"2022-05-13T00:57:35.757768Z","shell.execute_reply.started":"2022-05-13T00:57:35.724607Z","shell.execute_reply":"2022-05-13T00:57:35.757063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# src code modified and piggback\n\nclass Transformer(Module):\n    def __init__(self, d_model: int = 512, nhead: int = 8, num_encoder_layers: int = 6,\n                 num_decoder_layers: int = 6, dim_feedforward: int = 2048, dropout: float = 0.1,\n                 activation: Union[str, Callable[[Tensor], Tensor]] = F.relu,\n                 custom_encoder: Optional[Any] = None, custom_decoder: Optional[Any] = None,\n                 layer_norm_eps: float = 1e-5, batch_first: bool = False, norm_first: bool = False,\n                 device=None, dtype=None) -> None:\n#         factory_kwargs = {'device': device, 'dtype': dtype}\n#         print(factory_kwargs)\n        super(Transformer, self).__init__()\n\n        if custom_encoder is not None:\n            self.encoder = custom_encoder\n        else:\n            encoder_layer = TransformerEncoderLayer(d_model, nhead, dim_feedforward,\n                                                    dropout, device=device, dtype=dtype)\n            encoder_norm = LayerNorm(d_model, eps=layer_norm_eps, device=device, dtype=dtype)\n            self.encoder = TransformerEncoder(encoder_layer, num_encoder_layers, encoder_norm)\n\n        if custom_decoder is not None:\n            self.decoder = custom_decoder\n        else:\n            decoder_layer = TransformerDecoderLayer(d_model, nhead, dim_feedforward, dropout,\n                                                     device=device, dtype=dtype)                                                    \n            decoder_norm = LayerNorm(d_model, eps=layer_norm_eps, device=device, dtype=dtype)\n            self.decoder = TransformerDecoder(decoder_layer, num_decoder_layers, decoder_norm)\n\n        self._reset_parameters()\n\n        self.d_model = d_model\n        self.nhead = nhead\n\n        self.batch_first = batch_first\n        # Adding Context ####\n        self.norm = nn.LayerNorm(d_model)\n\n        self.feed_forward = nn.Sequential(\n            nn.Linear(2*d_model, dim_feedforward),\n            nn.ReLU(),\n            nn.Linear(dim_feedforward, d_model),\n        )\n\n        self.dropout = nn.Dropout(dropout)\n        ### End Add Context #######\n        \n#     def make_src_mask(self, src):\n#         src_mask = src.transpose(0, 1) == 105540\n\n#         # (N, src_len)\n#         return src_mask.to(self.device)\n    def forward(self, src: Tensor, tgt: Tensor, custD: Tensor, src_mask: Optional[Tensor] = None, tgt_mask: Optional[Tensor] = None,\n                memory_mask: Optional[Tensor] = None, src_key_padding_mask: Optional[Tensor] = None,\n                tgt_key_padding_mask: Optional[Tensor] = None, memory_key_padding_mask: Optional[Tensor] = None) -> Tensor:\n        \n        is_batched = src.dim() == 3\n        if not self.batch_first and src.size(1) != tgt.size(1) and is_batched:\n            raise RuntimeError(\"the batch number of src and tgt must be equal\")\n        elif self.batch_first and src.size(0) != tgt.size(0) and is_batched:\n            raise RuntimeError(\"the batch number of src and tgt must be equal\")\n\n        if src.size(-1) != self.d_model or tgt.size(-1) != self.d_model:\n            raise RuntimeError(\"the feature number of src and tgt must be equal to d_model\")\n        if src.size(0) != custD.size(0):\n            raise RuntimeError(\"the number of src and customer must be equal\")\n#         print(\"src\", src.shape, src)\n#         print(\"tgt\", tgt.shape, tgt)\n#         print(\"cust\", custD.shape, custD)\n        memory = self.encoder(src+custD, mask=src_mask, src_key_padding_mask=src_key_padding_mask)\n#         print(f\"memory shape {memory}\")\n        ### Add context ###############\n#         context_custExp = custD #context_cust.unsqueeze(1).repeat(1, self.max_length, 1)\n#         context_forward = self.feed_forward(torch.cat([memory, custD], dim=2))\n#         print(f\"context_forward shape {context_forward}\")\n#         context_out = self.dropout(self.norm(context_forward + custD)) # still give more importance to customer emb\n#         print(f\"context_out shape {context_out}\")\n#         context_out = memory #+ custD\n#         print(\"context_out\", context_out)\n        \n        #### End Context ###############\n        \n#         print(\"trgtP\", tgt_key_padding_mask)\n        output = self.decoder(tgt, memory, tgt_mask=tgt_mask, memory_mask=memory_mask,\n                              tgt_key_padding_mask=tgt_key_padding_mask,\n                              memory_key_padding_mask=memory_key_padding_mask)\n#         print(f\"output shape {output}\")\n        return output\n\n    @staticmethod\n    def generate_square_subsequent_mask(sz: int) -> Tensor:\n        r\"\"\"Generate a square mask for the sequence. The masked positions are filled with float('-inf').\n            Unmasked positions are filled with float(0.0).\n        \"\"\"\n        return torch.triu(torch.full((sz, sz), float('-inf')), diagonal=1)\n\n    def _reset_parameters(self):\n        r\"\"\"Initiate parameters in the transformer model.\"\"\"\n\n        for p in self.parameters():\n            if p.dim() > 1:\n                xavier_uniform_(p)","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:57:35.76277Z","iopub.execute_input":"2022-05-13T00:57:35.765181Z","iopub.status.idle":"2022-05-13T00:57:35.786365Z","shell.execute_reply.started":"2022-05-13T00:57:35.765138Z","shell.execute_reply":"2022-05-13T00:57:35.785549Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tqdm\nimport torch.optim as optim\nimport time, random\nimport numpy as np\n# import gc\n# gc.collect()\n# torch.cuda.empty_cache() \n\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n# device = torch.device(\"cpu\")\nsave_model = True\ntorch.manual_seed(163)\ntorch.backends.cudnn.benchmark = False\n# torch.set_deterministic(True)\nnp.random.seed(163)\nrandom.seed(163)\n\n# Training hyperparameters\nnum_epochs = 400\nlearning_rate = 1e-5\nbatch_size = 32\n\ntrain_iterator = BucketIterator(\n    train, \n    batch_size=batch_size,\n    device=\"cuda\"\n)\n\n# Model hyperparameters\nsrc_vocab_size = 105540# Number of Articles\ntrg_vocab_size = 105540# Number of Articles\nembedding_size = 256\nnum_heads = 8\nnum_encoder_layers = 8\nnum_decoder_layers = 8\ndropout_v = 0.10\nmax_len_s = 5 # Minimum Article Warm Start for Context\nmax_len_t = 10 # Maximum Article length to predict\nforward_expansion = 2048\nsrc_pad_idx = 105540 # english.vocab.stoi[\"<pad>\"]\n\n# Tensorboard to get nice loss plot\n# writer = SummaryWriter(\"runs/loss_plot\")\nstep = 0\n\n\nmodel = RecSTransformer(embedding_size,\n    src_vocab_size,\n    trg_vocab_size,\n    src_pad_idx,\n    num_heads,\n    num_encoder_layers,\n    num_decoder_layers,\n    forward_expansion,\n    dropout_v,\n    max_len_s,\n    max_len_t,\n    device).to(device)\n\n# last_model = wandb.restore('0-end-nan.pt', \n#                            run_path=\"mayankk-om-dev/HnmRGCNv2/3sbjjxyv\")\n\n# # use the \"name\" attribute of the returned object if your framework expects a filename, e.g. as in Keras\n# checkpoint = torch.load(last_model.name)\n\n# model.load_state_dict(checkpoint[\"state_dict\"])\n\n# wandb.watch(model)\n# print(\"Successfully Loaded Model\")\noptimizer = optim.Adam(model.parameters(), betas=(0.9, 0.98),eps=1e-09, lr=learning_rate)\n\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer, factor=0.1, patience=2, verbose=True\n)\n\n\n# optimizer = ScheduledOptim(\n#         optim.Adam(model.parameters(), betas=(0.9, 0.98), eps=1e-09),\n#         2, 256, 4000)\n\n\npad_idx = 105540\ncriterion = nn.CrossEntropyLoss(ignore_index=pad_idx)\n\nfor epoch in range(num_epochs):\n    print(f\"[Epoch {epoch} / {num_epochs}]\")\n\n    \n    model.train()\n    losses = []\n    with tqdm.tqdm(train_iterator) as tq:\n        for batch_idx, batch in enumerate(tq):\n#             print(batch)\n            # Get input and targets and get to cuda\n            customers = batch.c.to(device)\n            article_src = batch.asrc.to(device)\n            article_trg = batch.atrg.to(device)\n            customers = customers.view(1, customers.shape[0])\n#             print(article_trg[-10:, :])\n#             print(customers.shape)\n            # Forward prop\n            output = model(article_src, article_trg[:-1, :], customers)\n\n            # Output is of shape (trg_len, batch_size, output_dim) but Cross Entropy Loss\n            # doesn't take input in that form. For example if we have MNIST we want to have\n            # output to be: (N, 10) and targets just (N). Here we can view it in a similar\n            # way that we have output_words * batch_size that we want to send in into\n            # our cost function, so we need to do some reshapin.\n            # Let's also remove the start token while we're at it\n            output = output.reshape(-1, output.shape[2])\n#             print(output)\n            target = article_trg[1:].reshape(-1)\n#             print(\"-\"*100)\n#             print(target)\n#             print(\"#\"*100)\n            optimizer.zero_grad()\n\n            loss = criterion(output, target)\n#             print(\"*\"*10)\n#             print(loss.item())\n            losses.append(loss.item())\n\n            # Back prop\n            loss.backward()\n            # Clip to avoid exploding gradient issues, makes sure grads are\n            # within a healthy range\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=0.75, error_if_nonfinite=False)\n\n            # Gradient descent step\n            optimizer.step()\n#             optimizer.step_and_update_lr()\n            \n            tq.set_postfix({'loss': '%.03f' % loss.item()}, refresh=False)\n            wandb.log({\"train_loss_mb\": loss.item()})\n            # plot to tensorboard\n#             writer.add_scalar(\"Training loss\", loss, global_step=step)\n            step += 1\n\n        checkpoint = {\n                    \"state_dict\": model.state_dict(),\n                    \"optimizer\": optimizer.state_dict(),\n                }\n        f_ = f'{epoch}-end-{loss.item()}.pt'\n\n        torch.save(checkpoint, f_)\n        lm = f_\n#         wandb.save(f_, policy='now')\n#         if epoch >=5 and epoch % 5==0:\n#             for fnm in os.listdir():\n#                 if fnm.endswith('.pt'):\n#                     if fnm == lm:continue\n#                     os.remove(fnm)\n        \n        if step%5000:\n            mean_loss = sum(losses) / len(losses)\n            scheduler.step(mean_loss)\n            losses = []","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:57:35.787723Z","iopub.execute_input":"2022-05-13T00:57:35.788235Z","iopub.status.idle":"2022-05-13T00:58:19.319109Z","shell.execute_reply.started":"2022-05-13T00:57:35.788198Z","shell.execute_reply":"2022-05-13T00:58:19.317933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:58:19.320316Z","iopub.status.idle":"2022-05-13T00:58:19.321265Z","shell.execute_reply.started":"2022-05-13T00:58:19.321Z","shell.execute_reply":"2022-05-13T00:58:19.321025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache() ","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:58:19.3224Z","iopub.status.idle":"2022-05-13T00:58:19.322946Z","shell.execute_reply.started":"2022-05-13T00:58:19.322704Z","shell.execute_reply":"2022-05-13T00:58:19.322729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# transformer_model = Transformer(d_model=256,nhead=8, num_encoder_layers=6)\n# src = torch.rand((5,256,32, 100))\n# cust = torch.rand((5, 256 ,32,100 ))\n# tgt = torch.rand((10, 256, 32,100))\n# out = transformer_model(src, tgt, cust)","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:58:19.324335Z","iopub.status.idle":"2022-05-13T00:58:19.325114Z","shell.execute_reply.started":"2022-05-13T00:58:19.324856Z","shell.execute_reply":"2022-05-13T00:58:19.324881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch\n# import torch.nn as nn\n\n# # Sratch Block with New Architecutre\n# class SelfAttention(nn.Module):\n#     def __init__(self, embed_size, heads):\n#         super(SelfAttention, self).__init__()\n#         self.embed_size = embed_size\n#         self.heads = heads\n#         self.head_dim = embed_size // heads\n\n#         assert (\n#             self.head_dim * heads == embed_size\n#         ), \"Embedding size needs to be divisible by heads\"\n\n#         self.values = nn.Linear(self.head_dim, self.head_dim, bias=False)\n#         self.keys = nn.Linear(self.head_dim, self.head_dim, bias=False)\n#         self.queries = nn.Linear(self.head_dim, self.head_dim, bias=False)\n#         self.fc_out = nn.Linear(heads * self.head_dim, embed_size)\n\n#     def forward(self, values, keys, query, mask):\n#         # Get number of training examples\n#         N = query.shape[0]\n\n#         value_len, key_len, query_len = values.shape[1], keys.shape[1], query.shape[1]\n\n#         # Split the embedding into self.heads different pieces\n#         values = values.reshape(N, value_len, self.heads, self.head_dim)\n#         keys = keys.reshape(N, key_len, self.heads, self.head_dim)\n#         query = query.reshape(N, query_len, self.heads, self.head_dim)\n\n#         values = self.values(values)  # (N, value_len, heads, head_dim)\n#         keys = self.keys(keys)  # (N, key_len, heads, head_dim)\n#         queries = self.queries(query)  # (N, query_len, heads, heads_dim)\n\n#         # Einsum does matrix mult. for query*keys for each training example\n#         # with every other training example, don't be confused by einsum\n#         # it's just how I like doing matrix multiplication & bmm\n\n#         energy = torch.einsum(\"nqhd,nkhd->nhqk\", [queries, keys])\n#         # queries shape: (N, query_len, heads, heads_dim),\n#         # keys shape: (N, key_len, heads, heads_dim)\n#         # energy: (N, heads, query_len, key_len)\n\n#         # Mask padded indices so their weights become 0\n#         if mask is not None:\n#             energy = energy.masked_fill(mask == 0, float(\"-1e20\"))\n\n#         # Normalize energy values similarly to seq2seq + attention\n#         # so that they sum to 1. Also divide by scaling factor for\n#         # better stability\n#         attention = torch.softmax(energy / (self.embed_size ** (1 / 2)), dim=3)\n#         # attention shape: (N, heads, query_len, key_len)\n\n#         out = torch.einsum(\"nhql,nlhd->nqhd\", [attention, values]).reshape(\n#             N, query_len, self.heads * self.head_dim\n#         )\n#         # attention shape: (N, heads, query_len, key_len)\n#         # values shape: (N, value_len, heads, heads_dim)\n#         # out after matrix multiply: (N, query_len, heads, head_dim), then\n#         # we reshape and flatten the last two dimensions.\n\n#         out = self.fc_out(out)\n#         # Linear layer doesn't modify the shape, final shape will be\n#         # (N, query_len, embed_size)\n\n#         return out\n\n\n# class TransformerBlock(nn.Module):\n#     def __init__(self, embed_size, heads, dropout, forward_expansion):\n#         super(TransformerBlock, self).__init__()\n#         self.attention = SelfAttention(embed_size, heads)\n#         self.norm1 = nn.LayerNorm(embed_size)\n#         self.norm2 = nn.LayerNorm(embed_size)\n\n#         self.feed_forward = nn.Sequential(\n#             nn.Linear(embed_size, forward_expansion * embed_size),\n#             nn.ReLU(),\n#             nn.Linear(forward_expansion * embed_size, embed_size),\n#         )\n\n#         self.dropout = nn.Dropout(dropout)\n\n#     def forward(self, value, key, query, mask):\n#         attention = self.attention(value, key, query, mask)\n\n#         # Add skip connection, run through normalization and finally dropout\n#         x = self.dropout(self.norm1(attention + query))\n#         forward = self.feed_forward(x)\n#         out = self.dropout(self.norm2(forward + x))\n#         return out\n\n\n# class Encoder(nn.Module):\n#     def __init__(\n#         self,\n#         src_vocab_size, # article_size for me articles 1--n may only use partial vocab\n#         embed_size, #256\n#         num_layers,\n#         heads,\n#         device,\n#         forward_expansion,\n#         dropout,\n#         max_length,\n#     ):\n\n#         super(Encoder, self).__init__()\n#         self.embed_size = embed_size\n#         self.device = device\n#         weight_a = load_wandb(\"articleGNNEmb.pt\", \"mayankk-om-dev/HnmRGCNv2/2oooot9x\")\n#         weight_a = torch.cat([weight_a, torch.Tensor(1, 256)]) # included pad as last token\n#         self.article_embedding = nn.Embedding.from_pretrained(weight_a)\n        \n# #         assert weight_c.shape src_vocab_size, embed_size\n#         self.position_embedding = nn.Embedding(max_length, embed_size)\n\n#         self.layers = nn.ModuleList(\n#             [\n#                 TransformerBlock(\n#                     embed_size,\n#                     heads,\n#                     dropout=dropout,\n#                     forward_expansion=forward_expansion,\n#                 )\n#                 for _ in range(num_layers)\n#             ]\n#         )\n\n#         self.dropout = nn.Dropout(dropout)\n\n#     def forward(self, x, mask):\n#         N, seq_length = x.shape\n#         positions = torch.arange(0, seq_length).expand(N, seq_length).to(self.device)\n#         out = self.dropout(\n#             (self.article_embedding(x) + self.position_embedding(positions))\n#         )\n\n#         # In the Encoder the query, key, value are all the same, it's in the\n#         # decoder this will change. This might look a bit odd in this case.\n#         for layer in self.layers:\n#             out = layer(out, out, out, mask)\n\n#         return out\n\n\n# class DecoderBlock(nn.Module):\n#     def __init__(self, embed_size, heads, forward_expansion, dropout, device):\n#         super(DecoderBlock, self).__init__()\n#         self.norm = nn.LayerNorm(embed_size)\n#         self.attention = SelfAttention(embed_size, heads=heads)\n#         self.transformer_block = TransformerBlock(\n#             embed_size, heads, dropout, forward_expansion\n#         )\n#         self.dropout = nn.Dropout(dropout)\n\n#     def forward(self, x, value, key, src_mask, trg_mask):\n#         attention = self.attention(x, x, x, trg_mask)\n#         query = self.dropout(self.norm(attention + x))\n#         out = self.transformer_block(value, key, query, src_mask)\n#         return out\n\n\n# class Decoder(nn.Module):\n#     def __init__(\n#         self,\n#         trg_vocab_size,\n#         embed_size,\n#         num_layers,\n#         heads,\n#         forward_expansion,\n#         dropout,\n#         device,\n#         max_length,\n#     ):\n#         super(Decoder, self).__init__()\n#         self.device = device\n# #         self.word_embedding = nn.Embedding(trg_vocab_size, embed_size)\n#         weight_a = load_wandb(\"articleGNNEmb.pt\", \"mayankk-om-dev/HnmRGCNv2/2oooot9x\")\n#         weight_a = torch.cat([weight_a, torch.Tensor(1, 256)]) # included pad as last token\n#         self.article_embedding = nn.Embedding.from_pretrained(weight_a)\n        \n        \n#         self.position_embedding = nn.Embedding(max_length, embed_size)\n\n#         self.layers = nn.ModuleList(\n#             [\n#                 DecoderBlock(embed_size, heads, forward_expansion, dropout, device)\n#                 for _ in range(num_layers)\n#             ]\n#         )\n#         self.fc_out = nn.Linear(embed_size, trg_vocab_size)\n#         self.dropout = nn.Dropout(dropout)\n\n#     def forward(self, x, enc_out, src_mask, trg_mask):\n#         N, seq_length = x.shape\n#         positions = torch.arange(0, seq_length).expand(N, seq_length).to(self.device)\n#         x = self.dropout((self.article_embedding(x) + self.position_embedding(positions)))\n\n#         for layer in self.layers:\n#             x = layer(x, enc_out, enc_out, src_mask, trg_mask)\n\n#         out = self.fc_out(x)\n\n#         return out\n\n\n# class VTransformer(nn.Module):\n#     def __init__(\n#         self,\n#         src_vocab_size,\n#         trg_vocab_size,\n#         src_pad_idx,\n#         trg_pad_idx,\n#         embed_size=256,\n#         num_layers=6,\n#         forward_expansion=4,\n#         heads=8,\n#         dropout=0,\n#         device=\"cpu\",\n#         max_ls=5,\n#         max_lt=10,\n#     ):\n\n#         super(VTransformer, self).__init__()\n\n#         self.encoder = Encoder(\n#             src_vocab_size,\n#             embed_size,\n#             num_layers,\n#             heads,\n#             device,\n#             forward_expansion,\n#             dropout,\n#             max_ls,\n#         )\n\n#         self.decoder = Decoder(\n#             trg_vocab_size,\n#             embed_size,\n#             num_layers,\n#             heads,\n#             forward_expansion,\n#             dropout,\n#             device,\n#             max_lt,\n#         )\n#         # Adding Context ####\n# #         self.norm = nn.LayerNorm(embed_size)\n\n# #         self.feed_forward = nn.Sequential(\n# #             nn.Linear(2*embed_size, forward_expansion *2* embed_size),\n# #             nn.ReLU(),\n# #             nn.Linear(forward_expansion * 2* embed_size, embed_size),\n# #         )\n\n# #         self.dropout = nn.Dropout(dropout)\n#         weight_c = load_wandb(\"customerGNNEmb.pt\", \"mayankk-om-dev/HnmRGCNv2/2oooot9x\")\n#         self.context_embedding = nn.Embedding.from_pretrained(weight_c)\n        \n#         ### End Add Context #######\n#         self.src_pad_idx = src_pad_idx\n#         self.trg_pad_idx = trg_pad_idx\n#         self.device = device\n#         self.max_ls = max_ls\n#         self.max_lt = max_lt\n\n#     def make_src_mask(self, src):\n#         src_mask = (src != self.src_pad_idx).unsqueeze(1).unsqueeze(2)\n#         # (N, 1, 1, src_len)\n#         return src_mask.to(self.device)\n\n#     def make_trg_mask(self, trg):\n#         N, trg_len = trg.shape\n#         trg_mask = torch.tril(torch.ones((trg_len, trg_len))).expand(\n#             N, 1, trg_len, trg_len\n#         )\n\n#         return trg_mask.to(self.device)\n\n#     def forward(self, src, trg, customer): # how to pass customer\n#         src_mask = self.make_src_mask(src)\n#         trg_mask = self.make_trg_mask(trg)\n#         enc_src = self.encoder(src, src_mask)\n#         print(enc_src.shape)\n#         # customer emb N,\n# #         context_cust = self.customer_embedding(customer) # N, emb_size_c\n# #         context_custExp = context_cust.unsqueeze(1).repeat(1, self.max_length, 1)\n# #         # somehow reshape to N, query_len, emb_size_c (repeat)\n# #         context_forward = self.feed_forward(torch.concat([enc_src, context_custExp], dim=2))\n# #         context_out = self.dropout(self.norm(context_forward + context_cust)) # still give more importance to customer emb\n        \n#         embed_cust = self.context_embedding(cust)\n#         embed_cust = embed_cust.repeat(self.max_ls,1, 1) #N, seq_len, embs\n#         print(embed_cust.shape)\n#         context_out = embed_cust + enc_src\n        \n#         out = self.decoder(trg, context_out, src_mask, trg_mask)\n#         return out\n\n","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:58:19.326568Z","iopub.status.idle":"2022-05-13T00:58:19.327198Z","shell.execute_reply.started":"2022-05-13T00:58:19.326941Z","shell.execute_reply":"2022-05-13T00:58:19.326966Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import tqdm\n# import torch.optim as optim\n# import time, random\n# import numpy as np\n# # import gc\n# # gc.collect()\n# # torch.cuda.empty_cache() \n\n\n# device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n# # device = torch.device(\"cpu\")\n# save_model = True\n# torch.manual_seed(99)\n# torch.backends.cudnn.benchmark = False\n# # torch.set_deterministic(True)\n# np.random.seed(99)\n# random.seed(99)\n\n# # Training hyperparameters\n# num_epochs = 400\n# learning_rate = 3e-4\n# batch_size = 32\n\n# train_iterator = BucketIterator(\n#     train, \n#     batch_size=batch_size,\n#     device=\"cuda\"\n# )\n\n# # Model hyperparameters\n# src_vocab_size = 105540# Number of Articles\n# trg_vocab_size = 105540# Number of Articles\n# embedding_size = 256\n# num_heads = 8\n# num_encoder_layers = 4\n# num_decoder_layers = 8\n# dropout = 0.10\n# max_len_s = 5 # Minimum Article Warm Start for Context\n# max_len_t = 10 # Maximum Article length to predict\n# forward_expansion = 2048\n# src_pad_idx = 105540 # english.vocab.stoi[\"<pad>\"]\n# trg_pad_idx = 105540 \n# # Tensorboard to get nice loss plot\n# # writer = SummaryWriter(\"runs/loss_plot\")\n# step = 0\n\n\n# model = VTransformer(src_vocab_size,\n#         trg_vocab_size,\n#         src_pad_idx,\n#         trg_pad_idx,\n#         embed_size=256,\n#         num_layers=6,\n#         forward_expansion=4,\n#         heads=8,\n#         dropout=0.10,\n#         device=device,\n#         max_ls=5,\n#         max_lt=10,).to(device)\n\n# # last_model = wandb.restore('0-end-nan.pt', \n# #                            run_path=\"mayankk-om-dev/HnmRGCNv2/3sbjjxyv\")\n\n# # # use the \"name\" attribute of the returned object if your framework expects a filename, e.g. as in Keras\n# # checkpoint = torch.load(last_model.name)\n\n# # model.load_state_dict(checkpoint[\"state_dict\"])\n\n# # wandb.watch(model)\n# # print(\"Successfully Loaded Model\")\n# optimizer = optim.Adam(model.parameters(), betas=(0.9, 0.98),eps=1e-09, lr=learning_rate)\n\n# scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n#     optimizer, factor=0.1, patience=10, verbose=True\n# )\n\n\n# # optimizer = ScheduledOptim(\n# #         optim.Adam(model.parameters(), betas=(0.9, 0.98), eps=1e-09),\n# #         2, 256,4000 )\n\n\n# pad_idx = 105540\n# criterion = nn.CrossEntropyLoss(ignore_index=pad_idx)\n\n# for epoch in range(num_epochs):\n#     print(f\"[Epoch {epoch} / {num_epochs}]\")\n\n    \n#     model.train()\n#     losses = []\n#     with tqdm.tqdm(train_iterator) as tq:\n#         for batch_idx, batch in enumerate(tq):\n# #             print(batch)\n#             # Get input and targets and get to cuda\n#             customers = batch.c.to(device)\n#             article_src = batch.asrc.to(device)\n#             article_trg = batch.atrg.to(device)\n#             customers = customers.view(1, customers.shape[0])\n# #             print(article_trg[-10:, :])\n# #             print(customers.shape)\n#             # Forward prop\n#             output = model(article_src, article_trg[:-1, :], customers)\n\n#             # Output is of shape (trg_len, batch_size, output_dim) but Cross Entropy Loss\n#             # doesn't take input in that form. For example if we have MNIST we want to have\n#             # output to be: (N, 10) and targets just (N). Here we can view it in a similar\n#             # way that we have output_words * batch_size that we want to send in into\n#             # our cost function, so we need to do some reshapin.\n#             # Let's also remove the start token while we're at it\n#             output = output.reshape(-1, output.shape[2])\n#             print(output)\n#             target = article_trg[1:].reshape(-1)\n#             print(\"-\"*100)\n#             print(target)\n#             print(\"#\"*100)\n#             optimizer.zero_grad()\n\n#             loss = criterion(output, target)\n#             print(\"*\"*10)\n#             print(loss.item())\n#             losses.append(loss.item())\n\n#             # Back prop\n#             loss.backward()\n#             # Clip to avoid exploding gradient issues, makes sure grads are\n#             # within a healthy range\n#             torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=0.98, error_if_nonfinite=False)\n\n#             # Gradient descent step\n#             optimizer.step()\n# #             optimizer.step_and_update_lr()\n            \n#             tq.set_postfix({'loss': '%.03f' % loss.item()}, refresh=False)\n# #             wandb.log({\"train_loss_mb\": loss.item()})\n#             # plot to tensorboard\n# #             writer.add_scalar(\"Training loss\", loss, global_step=step)\n#             step += 1\n\n#         checkpoint = {\n#                     \"state_dict\": model.state_dict(),\n#                     \"optimizer\": optimizer.state_dict(),\n#                 }\n#         f_ = f'{epoch}-end-{loss.item()}.pt'\n\n#         torch.save(checkpoint, f_)\n#         lm = f_\n# #         wandb.save(f_, policy='now')\n#         if epoch >=5 and epoch % 5==0:\n#             for fnm in os.listdir():\n#                 if fnm.endswith('.pt'):\n#                     if fnm == lm:continue\n#                     os.remove(fnm)\n        \n#         mean_loss = sum(losses) / len(losses)\n#         scheduler.step(mean_loss)","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:58:19.328849Z","iopub.status.idle":"2022-05-13T00:58:19.329516Z","shell.execute_reply.started":"2022-05-13T00:58:19.329277Z","shell.execute_reply":"2022-05-13T00:58:19.329301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''A wrapper class for scheduled optimizer '''\nimport numpy as np\n\nclass ScheduledOptim():\n    '''A simple wrapper class for learning rate scheduling'''\n\n    def __init__(self, optimizer, lr_mul, d_model, n_warmup_steps):\n        self._optimizer = optimizer\n        self.lr_mul = lr_mul\n        self.d_model = d_model\n        self.n_warmup_steps = n_warmup_steps\n        self.n_steps = 0\n\n\n    def step_and_update_lr(self):\n        \"Step with the inner optimizer\"\n        self._update_learning_rate()\n        self._optimizer.step()\n\n\n    def zero_grad(self):\n        \"Zero out the gradients with the inner optimizer\"\n        self._optimizer.zero_grad()\n\n\n    def _get_lr_scale(self):\n        d_model = self.d_model\n        n_steps, n_warmup_steps = self.n_steps, self.n_warmup_steps\n        return (d_model ** -0.5) * min(n_steps ** (-0.5), n_steps * n_warmup_steps ** (-1.5))\n\n\n    def _update_learning_rate(self):\n        ''' Learning rate scheduling per step '''\n\n        self.n_steps += 1\n        lr = self.lr_mul * self._get_lr_scale()\n\n        for param_group in self._optimizer.param_groups:\n            param_group['lr'] = lr","metadata":{"execution":{"iopub.status.busy":"2022-05-13T00:58:19.330784Z","iopub.status.idle":"2022-05-13T00:58:19.331462Z","shell.execute_reply.started":"2022-05-13T00:58:19.331221Z","shell.execute_reply":"2022-05-13T00:58:19.331245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}