{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# TorchMetrics for Tile\nThis notebook presents a way of calculating the metric for the `tile:xla` collection in torch  \nThe metric expects the following inputs\n1. predicted_ordered_configs of shape `(bs, num_config)`\n2. config_runtime of shape `(bs, num_config)`\n3. config_attn_mask of shape `(bs, num_config)` (1, where the batch element has a valid configuration, 0 when not)\n\n\nThe attention mask is expected since not all models have the same number of configurations","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd \nfrom pathlib import Path\nimport torch\nfrom torch.utils.data import Dataset\nfrom torch.nn import functional as F\nfrom torch.nn.utils.rnn import pad_sequence\nimport torchmetrics as tm\nfrom dataclasses import dataclass\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-09-03T11:34:18.477842Z","iopub.execute_input":"2023-09-03T11:34:18.478798Z","iopub.status.idle":"2023-09-03T11:34:26.737725Z","shell.execute_reply.started":"2023-09-03T11:34:18.478750Z","shell.execute_reply":"2023-09-03T11:34:26.735840Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 1. Metric for the collection `tile:xla`\n- This metric is used specifically for the collection `tile:xla`.\n\n- `(1-slowdown)` inccured of the top-K predictions is used to reflect how much slower the top-K configurations predicted by the model is from the actual fastest configurations.\n\n- The metric can be formulated as follows:\n$$1 - \\left( \\frac{\\text{The best runtime of the top-k predictions}}{\\text{The best runtime of all configurations}} - 1 \\right) = 2 - \\frac{\\min_{i \\in K} y_i}{\\min_{i \\in A} y_i}$$\n Where K is the top-K predictions, A is all configurations of the given graph from the dataset collection, and y is the measured execution time.","metadata":{}},{"cell_type":"code","source":"DATA_DIR = \"../input/predict-ai-model-runtime/npz_all/npz\"","metadata":{"execution":{"iopub.status.busy":"2023-09-03T11:34:26.744697Z","iopub.execute_input":"2023-09-03T11:34:26.745561Z","iopub.status.idle":"2023-09-03T11:34:26.752171Z","shell.execute_reply.started":"2023-09-03T11:34:26.745480Z","shell.execute_reply":"2023-09-03T11:34:26.750167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tile_df = pd.DataFrame({'paths': [elem for elem in (Path(DATA_DIR) / 'tile').rglob(\"*\") if elem.is_file()]}).assign(\n    split=lambda df: df.paths.apply(lambda x: x.parent.name),\n    configuration=lambda df: df.paths.apply(lambda x: x.parent.parent.name),\n    extra=lambda df: df.paths.apply(lambda x: x.parent.parent.parent.name),\n    model_name=lambda df: df.paths.apply(lambda x: x.stem),\n    collection=lambda df: df.extra + ':' + df.configuration ,\n    paths = lambda df: df.paths.apply(lambda x: str(x))\n)\ntile_df","metadata":{"execution":{"iopub.status.busy":"2023-09-03T11:34:26.754521Z","iopub.execute_input":"2023-09-03T11:34:26.754919Z","iopub.status.idle":"2023-09-03T11:34:31.204914Z","shell.execute_reply.started":"2023-09-03T11:34:26.754887Z","shell.execute_reply":"2023-09-03T11:34:31.203617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Element Dataset with adjacency matrix","metadata":{}},{"cell_type":"code","source":"def edges_adjacency(edges: torch.Tensor) -> torch.Tensor:\n    adj = torch.zeros((edges.max() + 1, edges.max() + 1))\n    adj[edges[0], edges[1]] = 1\n    return adj","metadata":{"execution":{"iopub.status.busy":"2023-09-03T11:34:31.209459Z","iopub.execute_input":"2023-09-03T11:34:31.209883Z","iopub.status.idle":"2023-09-03T11:34:31.216311Z","shell.execute_reply.started":"2023-09-03T11:34:31.209848Z","shell.execute_reply":"2023-09-03T11:34:31.215138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def tile_loader(path):\n    tile_dict =  dict(np.load(path))\n    tile_dict = {k: torch.from_numpy(v) for k, v in tile_dict.items()}\n    tile_dict['edges_adjecency'] = edges_adjacency(tile_dict.pop('edge_index'))\n    return tile_dict","metadata":{"execution":{"iopub.status.busy":"2023-09-03T11:34:31.217549Z","iopub.execute_input":"2023-09-03T11:34:31.217913Z","iopub.status.idle":"2023-09-03T11:34:31.233335Z","shell.execute_reply.started":"2023-09-03T11:34:31.217884Z","shell.execute_reply":"2023-09-03T11:34:31.232150Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TileDataset(Dataset):\n    \n    def __init__(self, df):\n        self.df = df\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        tile_dict = tile_loader(self.df.paths[idx])\n        return tile_dict","metadata":{"execution":{"iopub.status.busy":"2023-09-03T11:34:31.234631Z","iopub.execute_input":"2023-09-03T11:34:31.235723Z","iopub.status.idle":"2023-09-03T11:34:31.249171Z","shell.execute_reply.started":"2023-09-03T11:34:31.235685Z","shell.execute_reply":"2023-09-03T11:34:31.248093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Collator with padding and attention masks","metadata":{}},{"cell_type":"code","source":"def pad_edge_adjacency(edges_adjacency_list):\n    max_len = max([elem.shape[0] for elem in edges_adjacency_list])\n    return torch.stack([F.pad(elem, (0, max_len-elem.shape[0], 0, max_len-elem.shape[0]), value=0) for elem in edges_adjacency_list], dim=0)\n\n@dataclass\nclass TileCollator:\n    pad_to_multiple_of: int = 64\n    targets:bool = True\n    \n    def __call__(self, batch):\n        output = {}\n        max_node_len = max([elem['node_opcode'].shape[0] for elem in batch])\n        node_pad_amount = self.pad_to_multiple_of - max_node_len % max(self.pad_to_multiple_of, 1)\n        output['node_opcode'] = F.pad(pad_sequence([elem['node_opcode'] for elem in batch], batch_first=True),\n                                      (0, node_pad_amount), value=0).long()\n        output['node_feat'] = F.pad(pad_sequence([elem['node_feat'] for elem in batch], batch_first=True),\n                                    (0,0,0, node_pad_amount), value=0)\n        output['edges_adjecency'] = F.pad(pad_edge_adjacency([elem['edges_adjecency'] for elem in batch]),\n                                          (0, node_pad_amount, 0, node_pad_amount), value=0)\n        output['node_attn_mask'] = F.pad(pad_sequence([torch.ones(len(elem['node_opcode'])) for elem in batch], batch_first=True),\n                                         (0, node_pad_amount), value=0)\n        max_config_len = max([elem['config_feat'].shape[0] for elem in batch])\n        config_pad_amount = self.pad_to_multiple_of - max_config_len % max(self.pad_to_multiple_of, 1)\n        \n        output['config_feat'] = F.pad(pad_sequence([elem['config_feat'] for elem in batch], batch_first=True),\n                                          (0, 0, 0, config_pad_amount), value=0)\n                                      \n        output['config_attn_mask'] = F.pad(pad_sequence([torch.ones(len(elem['config_feat']))  for elem in batch], batch_first=True),\n                                           (0, config_pad_amount), value=0)\n        \n        if self.targets:\n            output['config_runtime'] = F.pad(pad_sequence([elem['config_runtime'] / elem['config_runtime_normalizers'] for elem in batch], batch_first=True),\n                                             (0, config_pad_amount), value=0)\n        return output","metadata":{"execution":{"iopub.status.busy":"2023-09-03T11:34:31.251243Z","iopub.execute_input":"2023-09-03T11:34:31.252058Z","iopub.status.idle":"2023-09-03T11:34:31.274930Z","shell.execute_reply.started":"2023-09-03T11:34:31.252011Z","shell.execute_reply":"2023-09-03T11:34:31.273340Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tile_dataset = TileDataset(tile_df)\nbatch = [tile_dataset[i] for i in range(8)]\ncollate_fn = TileCollator(64)\nbatch = collate_fn(batch)","metadata":{"execution":{"iopub.status.busy":"2023-09-03T11:34:31.276783Z","iopub.execute_input":"2023-09-03T11:34:31.278063Z","iopub.status.idle":"2023-09-03T11:34:31.339404Z","shell.execute_reply.started":"2023-09-03T11:34:31.278011Z","shell.execute_reply":"2023-09-03T11:34:31.338103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for k,v in batch.items():\n    print(k, v.shape, v.dtype)","metadata":{"execution":{"iopub.status.busy":"2023-09-03T11:34:31.340794Z","iopub.execute_input":"2023-09-03T11:34:31.341258Z","iopub.status.idle":"2023-09-03T11:34:31.351572Z","shell.execute_reply.started":"2023-09-03T11:34:31.341221Z","shell.execute_reply":"2023-09-03T11:34:31.349931Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Tile Top K Metric\n1. Select the best runtimes configuration for each element in the target\n2. Select the best K indices from the predicted configurations tensor\n3. Select the best runtime for those indices\n4. Compare best runtimes vs predicted runtimes","metadata":{}},{"cell_type":"code","source":"class TileTopK(tm.Metric):\n    \n    def __init__(self, k:int=5) -> None:\n        super().__init__()\n        self.add_state(\"runtimes\", default=[], dist_reduce_fx=None)\n        self.k = k\n        \n    def update(self, preds: torch.Tensor, target: torch.Tensor, config_attn_mask:torch.Tensor) -> None:\n        best_runtimes = torch.where(config_attn_mask==1, target, torch.tensor(float('inf'))).min(1).values\n        masked_preds = torch.where(config_attn_mask==1, preds, torch.tensor(float('inf')))\n        pred_bottomk_indices = torch.topk(masked_preds, k=self.k, largest=False).indices\n        bs = preds.shape[0]\n        bottom_k_positions = torch.stack([torch.arange(bs).repeat_interleave(self.k).to(preds.device), pred_bottomk_indices.view(-1)])\n        predicted_runtimes = target[bottom_k_positions[0], bottom_k_positions[1]].view(bs,self.k)\n        best_predicted_runtimes = predicted_runtimes.min(1).values\n        self.runtimes.append(best_predicted_runtimes/ best_runtimes)\n        \n    def compute(self) -> torch.Tensor:\n        return (2-torch.cat(self.runtimes)).mean()\n","metadata":{"execution":{"iopub.status.busy":"2023-09-03T11:34:31.353950Z","iopub.execute_input":"2023-09-03T11:34:31.354676Z","iopub.status.idle":"2023-09-03T11:34:31.369515Z","shell.execute_reply.started":"2023-09-03T11:34:31.354630Z","shell.execute_reply":"2023-09-03T11:34:31.368114Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metric = TileTopK()\nmetric.update(batch['config_runtime'], batch['config_runtime'], batch['config_attn_mask'])\nmetric.compute()","metadata":{"execution":{"iopub.status.busy":"2023-09-03T11:34:31.371453Z","iopub.execute_input":"2023-09-03T11:34:31.371922Z","iopub.status.idle":"2023-09-03T11:34:31.396191Z","shell.execute_reply.started":"2023-09-03T11:34:31.371884Z","shell.execute_reply":"2023-09-03T11:34:31.394570Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metric.reset()\nnoise = torch.rand_like(batch['config_attn_mask'])\nmetric.update(noise, batch['config_runtime'], batch['config_attn_mask'])\nmetric.compute()","metadata":{"execution":{"iopub.status.busy":"2023-09-03T11:34:31.398294Z","iopub.execute_input":"2023-09-03T11:34:31.398766Z","iopub.status.idle":"2023-09-03T11:34:31.412846Z","shell.execute_reply.started":"2023-09-03T11:34:31.398727Z","shell.execute_reply":"2023-09-03T11:34:31.411569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}],"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"}}