{"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":"code","source":"from typing import Any, Tuple\nimport torch\nimport torchmetrics as tm\nfrom pathlib import Path\nfrom scipy.stats import kendalltau\nimport numpy as np","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-09-13T22:45:41.984342Z","iopub.execute_input":"2023-09-13T22:45:41.984776Z","iopub.status.idle":"2023-09-13T22:45:41.990865Z","shell.execute_reply.started":"2023-09-13T22:45:41.984741Z","shell.execute_reply":"2023-09-13T22:45:41.989913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def calculate_concordant_discordant(true_sequence:torch.Tensor, pred_sequence:torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:\n    \"\"\"\n    Calculate the number of concordant and discordant pairs\n    Args:\n        true_sequence: Tensor of shape (bs, seq_len) with the true sequence\n        pred_sequence: Tensor of shape (bs, seq_len) with the predicted sequence\n    Returns:\n        concordant: Tensor of shape (bs,) with the number of concordant pairs\n        discordant: Tensor of shape (bs,) with the number of discordant pairs\n    \"\"\"\n    num_configs = true_sequence.shape[1]\n    tril_mask = torch.ones((num_configs, num_configs), device=true_sequence.device).tril(diagonal=-1)\n    true_diff = (true_sequence.unsqueeze(-1) - true_sequence.unsqueeze(1))\n    pred_diff = pred_sequence.unsqueeze(-1) - pred_sequence.unsqueeze(1)\n    concordant = ((true_diff * pred_diff > 0).float() * tril_mask).sum(dim=[1,2])\n    discordant = ((true_diff * pred_diff < 0).float() * tril_mask).sum(dim=[1,2])\n    return concordant, discordant\n\nclass KendallTau(tm.Metric):\n    \n    higher_is_better = True\n    \n    def __init__(self, eps:float=1e-6, **kwargs: Any) -> None:\n        super().__init__(**kwargs)\n        self.add_state(\"concordant\", default=[], dist_reduce_fx=None)\n        self.add_state(\"discordant\", default=[], dist_reduce_fx=None)\n        self.eps = eps\n        \n    def update(self, true_sequence:torch.Tensor, pred_sequence:torch.Tensor):\n        concordant, discordant = calculate_concordant_discordant(true_sequence, pred_sequence)\n        self.concordant.append(concordant)\n        self.discordant.append(discordant)\n        \n    def kendall_tau(self):\n        concordant = torch.cat(self.concordant)\n        discordant = torch.cat(self.discordant)\n        kendall_tau = (concordant - discordant) / (concordant + discordant + self.eps)\n        return kendall_tau\n        \n    def compute(self) -> torch.Tensor:\n        kendall_tau = self.kendall_tau()\n        return kendall_tau.mean()","metadata":{"execution":{"iopub.status.busy":"2023-09-13T22:45:42.000601Z","iopub.execute_input":"2023-09-13T22:45:42.001009Z","iopub.status.idle":"2023-09-13T22:45:42.025025Z","shell.execute_reply.started":"2023-09-13T22:45:42.000972Z","shell.execute_reply":"2023-09-13T22:45:42.023948Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_paths = list(Path(\"/kaggle/input/predict-ai-model-runtime/npz_all/npz/layout/nlp/default/train\").rglob(\"*.npz\"))","metadata":{"execution":{"iopub.status.busy":"2023-09-13T22:45:42.030670Z","iopub.execute_input":"2023-09-13T22:45:42.031914Z","iopub.status.idle":"2023-09-13T22:45:42.053019Z","shell.execute_reply.started":"2023-09-13T22:45:42.031867Z","shell.execute_reply":"2023-09-13T22:45:42.051813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"first_runtimes = [np.load(path)['config_runtime'] for path in all_paths[:50]]\nshortest_runtime = min(*[len(elem) for elem in first_runtimes], 500)\nfirst_runtimes = [elem[:shortest_runtime] for elem in first_runtimes]","metadata":{"execution":{"iopub.status.busy":"2023-09-13T22:45:42.057712Z","iopub.execute_input":"2023-09-13T22:45:42.060034Z","iopub.status.idle":"2023-09-13T22:45:42.353404Z","shell.execute_reply.started":"2023-09-13T22:45:42.060001Z","shell.execute_reply":"2023-09-13T22:45:42.352382Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"first_runtimes_np = np.stack(first_runtimes)\nruntime_noise = np.random.random(size = first_runtimes_np.shape)\nfirst_runtimes_np.shape, runtime_noise.shape","metadata":{"execution":{"iopub.status.busy":"2023-09-13T22:45:42.360337Z","iopub.execute_input":"2023-09-13T22:45:42.362794Z","iopub.status.idle":"2023-09-13T22:45:42.374175Z","shell.execute_reply.started":"2023-09-13T22:45:42.362759Z","shell.execute_reply":"2023-09-13T22:45:42.373164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"first_runtimes_np","metadata":{"execution":{"iopub.status.busy":"2023-09-13T22:45:42.378631Z","iopub.execute_input":"2023-09-13T22:45:42.380909Z","iopub.status.idle":"2023-09-13T22:45:42.390845Z","shell.execute_reply.started":"2023-09-13T22:45:42.380872Z","shell.execute_reply":"2023-09-13T22:45:42.389980Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nfor elem, noise in zip(first_runtimes_np, runtime_noise):\n    print(kendalltau(elem, noise))","metadata":{"execution":{"iopub.status.busy":"2023-09-13T22:45:42.395476Z","iopub.execute_input":"2023-09-13T22:45:42.397829Z","iopub.status.idle":"2023-09-13T22:45:42.437926Z","shell.execute_reply.started":"2023-09-13T22:45:42.397797Z","shell.execute_reply":"2023-09-13T22:45:42.437068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'","metadata":{"execution":{"iopub.status.busy":"2023-09-13T22:45:42.442859Z","iopub.execute_input":"2023-09-13T22:45:42.445206Z","iopub.status.idle":"2023-09-13T22:45:42.451235Z","shell.execute_reply.started":"2023-09-13T22:45:42.445171Z","shell.execute_reply":"2023-09-13T22:45:42.450300Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kendalltau_metric = KendallTau()\nfirst_runtimes_torch = torch.from_numpy(first_runtimes_np).to(device)\nruntime_noise_torch = torch.from_numpy(runtime_noise).to(device)","metadata":{"execution":{"iopub.status.busy":"2023-09-13T22:45:42.455649Z","iopub.execute_input":"2023-09-13T22:45:42.457741Z","iopub.status.idle":"2023-09-13T22:45:42.467300Z","shell.execute_reply.started":"2023-09-13T22:45:42.456406Z","shell.execute_reply":"2023-09-13T22:45:42.466186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nkendalltau_metric.update(first_runtimes_torch, runtime_noise_torch)\nkendalltau_metric.kendall_tau()","metadata":{"execution":{"iopub.status.busy":"2023-09-13T22:45:42.471019Z","iopub.execute_input":"2023-09-13T22:45:42.471735Z","iopub.status.idle":"2023-09-13T22:45:42.493879Z","shell.execute_reply.started":"2023-09-13T22:45:42.471701Z","shell.execute_reply":"2023-09-13T22:45:42.493065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time \nconcordant, discordant = calculate_concordant_discordant(first_runtimes_torch, runtime_noise_torch)\n(concordant - discordant) / (concordant + discordant)","metadata":{"execution":{"iopub.status.busy":"2023-09-13T22:45:42.499527Z","iopub.execute_input":"2023-09-13T22:45:42.502131Z","iopub.status.idle":"2023-09-13T22:45:42.521475Z","shell.execute_reply.started":"2023-09-13T22:45:42.502098Z","shell.execute_reply":"2023-09-13T22:45:42.520380Z"},"trusted":true},"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"}}