{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"},"papermill":{"default_parameters":{},"duration":9735.485807,"end_time":"2023-09-01T19:32:52.176348","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2023-09-01T16:50:36.690541","version":"2.4.0"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":58266,"databundleVersionId":6509002,"sourceType":"competition"}],"dockerImageVersionId":30528,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# 📦 Install Required Packages\n# In this cell, we're installing essential Python packages using pip.\n# We are adding 'torch-geometric' and 'torch-scatter' packages to the environment.\n# These packages are crucial for graph-based machine learning tasks.\n\n!pip install torch-geometric torch-scatter\n","metadata":{"_kg_hide-output":true,"papermill":{"duration":264.245878,"end_time":"2023-09-01T16:55:11.276628","exception":false,"start_time":"2023-09-01T16:50:47.03075","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-09-13T21:02:27.028439Z","iopub.execute_input":"2023-09-13T21:02:27.032244Z","iopub.status.idle":"2023-09-13T21:02:41.339945Z","shell.execute_reply.started":"2023-09-13T21:02:27.032177Z","shell.execute_reply":"2023-09-13T21:02:41.33848Z"},"trusted":true},"execution_count":1,"outputs":[{"name":"stdout","text":"Requirement already satisfied: torch-geometric in /opt/conda/lib/python3.10/site-packages (2.3.1)\nRequirement already satisfied: torch-scatter in /opt/conda/lib/python3.10/site-packages (2.1.1)\nRequirement already satisfied: tqdm in /opt/conda/lib/python3.10/site-packages (from torch-geometric) (4.65.0)\nRequirement already satisfied: numpy in /opt/conda/lib/python3.10/site-packages (from torch-geometric) (1.23.5)\nRequirement already satisfied: scipy in /opt/conda/lib/python3.10/site-packages (from torch-geometric) (1.11.1)\nRequirement already satisfied: jinja2 in /opt/conda/lib/python3.10/site-packages (from torch-geometric) (3.1.2)\nRequirement already satisfied: requests in /opt/conda/lib/python3.10/site-packages (from torch-geometric) (2.31.0)\nRequirement already satisfied: pyparsing in /opt/conda/lib/python3.10/site-packages (from torch-geometric) (3.0.9)\nRequirement already satisfied: scikit-learn in /opt/conda/lib/python3.10/site-packages (from torch-geometric) (1.2.2)\nRequirement already satisfied: psutil>=5.8.0 in /opt/conda/lib/python3.10/site-packages (from torch-geometric) (5.9.3)\nRequirement already satisfied: MarkupSafe>=2.0 in /opt/conda/lib/python3.10/site-packages (from jinja2->torch-geometric) (2.1.3)\nRequirement already satisfied: charset-normalizer<4,>=2 in /opt/conda/lib/python3.10/site-packages (from requests->torch-geometric) (3.1.0)\nRequirement already satisfied: idna<4,>=2.5 in /opt/conda/lib/python3.10/site-packages (from requests->torch-geometric) (3.4)\nRequirement already satisfied: urllib3<3,>=1.21.1 in /opt/conda/lib/python3.10/site-packages (from requests->torch-geometric) (1.26.15)\nRequirement already satisfied: certifi>=2017.4.17 in /opt/conda/lib/python3.10/site-packages (from requests->torch-geometric) (2023.5.7)\nRequirement already satisfied: joblib>=1.1.1 in /opt/conda/lib/python3.10/site-packages (from scikit-learn->torch-geometric) (1.2.0)\nRequirement already satisfied: threadpoolctl>=2.0.0 in /opt/conda/lib/python3.10/site-packages (from scikit-learn->torch-geometric) (3.1.0)\n","output_type":"stream"}]},{"cell_type":"code","source":"# 📚 Import Necessary Libraries\n# In this cell, we import essential Python libraries to support various tasks in our notebook.\n# We use libraries like NumPy, Pandas, tqdm, scikit-learn, PyTorch, and others.\n# The 'device' variable is set to 'cuda' if a GPU is available; otherwise, it defaults to 'cpu'.\n\nimport numpy as np  # 🧮 NumPy for numerical computations\nimport pandas as pd  # 🐼 Pandas for data manipulation\nimport os  # 📂 Operating system-related functions\nfrom tqdm import tqdm  # 🔄 tqdm for progress bar visualization\n\nimport sklearn  # 🧬 scikit-learn for machine learning utilities\nimport sklearn.model_selection  # 📊 scikit-learn's model selection module\nimport torch  # 🔥 PyTorch for deep learning\nfrom torch import nn  # 🧠 PyTorch's neural network module\nfrom torch import Tensor  # 🚀 PyTorch's Tensor data type\nfrom torch_geometric.nn import GCNConv  # 📊 Graph Convolutional Network layer\nfrom torch_geometric.datasets import Planetoid  # 🌍 PyTorch Geometric dataset for graph data\nfrom torch.utils.data import DataLoader, Dataset  # 📦 PyTorch data loading utilities\nfrom timm.scheduler import CosineLRScheduler  # 📈 Learning rate scheduler\nimport matplotlib.pyplot as plt  # 📊 Matplotlib for plotting\n\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'  # ⚙️ Determine if CUDA (GPU) is available\n\ndevice","metadata":{"papermill":{"duration":4.819384,"end_time":"2023-09-01T16:55:16.104784","exception":false,"start_time":"2023-09-01T16:55:11.2854","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-09-13T21:08:30.630681Z","iopub.execute_input":"2023-09-13T21:08:30.631136Z","iopub.status.idle":"2023-09-13T21:08:34.161213Z","shell.execute_reply.started":"2023-09-13T21:08:30.631078Z","shell.execute_reply":"2023-09-13T21:08:34.160166Z"},"trusted":true},"execution_count":1,"outputs":[{"name":"stderr","text":"/opt/conda/lib/python3.10/site-packages/scipy/__init__.py:146: UserWarning: A NumPy version >=1.16.5 and <1.23.0 is required for this version of SciPy (detected version 1.23.5\n  warnings.warn(f\"A NumPy version >={np_minversion} and <{np_maxversion}\"\n","output_type":"stream"},{"execution_count":1,"output_type":"execute_result","data":{"text/plain":"'cuda'"},"metadata":{}}]},{"cell_type":"code","source":"# 📁 Define a Function to Load DataFrames\n# This function loads data stored in different splits (train, valid, test) from a specified directory.\n# It reads files in the directory, extracts data using NumPy, and organizes it into DataFrames.\n\ndef load_df(directory, limit=-1):\n    splits = [\"train\", \"valid\", \"test\"]  # 🔄 List of data splits\n    dfs = dict()  # 📊 Dictionary to store DataFrames for each split\n    \n    for split in splits:\n        path = os.path.join(directory, split)  # 📂 Define the path to the split's directory\n        files = os.listdir(path)  # 🗂 Get a list of files in the split's directory\n        list_df = []  # 📄 List to store data dictionaries\n        \n        if limit >= 0:\n            files = files[0:limit]  # load a small number of data points when appropriate\n        \n        for file in files:\n            d = dict(np.load(os.path.join(path, file)))  # 📦 Load data using NumPy\n            d['file'] = file  # 📄 Include the file name in the data dictionary\n            list_df.append(d)  # 🧾 Append the data dictionary to the list\n        dfs[split] = pd.DataFrame.from_dict(list_df)  # 🐼 Create a DataFrame from the list of data dictionaries and store it in the dictionary\n    return dfs","metadata":{"papermill":{"duration":0.020227,"end_time":"2023-09-01T16:55:16.152594","exception":false,"start_time":"2023-09-01T16:55:16.132367","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-09-13T21:08:37.735715Z","iopub.execute_input":"2023-09-13T21:08:37.736557Z","iopub.status.idle":"2023-09-13T21:08:37.747792Z","shell.execute_reply.started":"2023-09-13T21:08:37.736506Z","shell.execute_reply":"2023-09-13T21:08:37.746231Z"},"trusted":true},"execution_count":2,"outputs":[]},{"cell_type":"code","source":"# 📄 Load data using the defined function and store it in the 'tile_xla' variable\ntile_xla = load_df(\"/kaggle/input/predict-ai-model-runtime/npz_all/npz/tile/xla/\")","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"layout_xla_default = load_df(\"/kaggle/input/predict-ai-model-runtime/npz_all/npz/layout/xla/default\", limit=1)","metadata":{"execution":{"iopub.status.busy":"2023-09-13T21:08:41.020414Z","iopub.execute_input":"2023-09-13T21:08:41.020835Z","iopub.status.idle":"2023-09-13T21:08:41.385869Z","shell.execute_reply.started":"2023-09-13T21:08:41.020803Z","shell.execute_reply":"2023-09-13T21:08:41.384668Z"},"trusted":true},"execution_count":3,"outputs":[]},{"cell_type":"code","source":"layout_xla_default","metadata":{"execution":{"iopub.status.busy":"2023-09-13T21:08:45.264422Z","iopub.execute_input":"2023-09-13T21:08:45.264897Z","iopub.status.idle":"2023-09-13T21:08:51.981545Z","shell.execute_reply.started":"2023-09-13T21:08:45.264862Z","shell.execute_reply":"2023-09-13T21:08:51.980206Z"},"trusted":true},"execution_count":4,"outputs":[{"execution_count":4,"output_type":"execute_result","data":{"text/plain":"{'train':                                            node_feat  \\\n 0  [[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0,...   \n \n                                          node_opcode  \\\n 0  [63, 63, 2, 63, 63, 57, 63, 63, 2, 63, 63, 2, ...   \n \n                                           edge_index  \\\n 0  [[2, 0], [2, 1], [5, 3], [5, 4], [8, 6], [8, 7...   \n \n                                     node_config_feat  \\\n 0  [[[3.0, 0.0, -1.0, -1.0, -1.0, -1.0, 0.0, 3.0,...   \n \n                                      node_config_ids  \\\n 0  [1001, 1034, 1062, 1087, 1113, 1140, 1165, 119...   \n \n                                       config_runtime  node_splits  \\\n 0  [73585984, 75217964, 73773521, 73779258, 73778...  [[0, 5605]]   \n \n                     file  \n 0  resnet50.8x8.fp16.npz  ,\n 'valid':                                            node_feat  \\\n 0  [[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0,...   \n \n                                          node_opcode  \\\n 0  [63, 63, 2, 63, 63, 57, 63, 63, 2, 63, 63, 2, ...   \n \n                                           edge_index  \\\n 0  [[2, 0], [2, 1], [5, 3], [5, 4], [8, 6], [8, 7...   \n \n                                     node_config_feat  \\\n 0  [[[0.0, 3.0, -1.0, -1.0, -1.0, -1.0, 0.0, 3.0,...   \n \n                                      node_config_ids  \\\n 0  [1010, 1044, 1072, 1097, 1123, 1150, 1175, 120...   \n \n                                       config_runtime  node_splits  \\\n 0  [238200912, 238212420, 238190565, 238209920, 2...  [[0, 5673]]   \n \n                     file  \n 0  resnet50.4x4.fp16.npz  ,\n 'test':                                            node_feat  \\\n 0  [[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0,...   \n \n                                          node_opcode  \\\n 0  [63, 63, 2, 63, 63, 2, 63, 63, 2, 63, 63, 2, 6...   \n \n                                           edge_index  \\\n 0  [[2, 0], [2, 1], [5, 3], [5, 4], [8, 6], [8, 7...   \n \n                                     node_config_feat  \\\n 0  [[[1.0, 0.0, -1.0, -1.0, -1.0, -1.0, 0.0, 1.0,...   \n \n                                      node_config_ids  \\\n 0  [107, 114, 121, 133, 135, 140, 142, 144, 146, ...   \n \n                                       config_runtime           node_splits  \\\n 0  [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, ...  [[0, 382, 432, 490]]   \n \n                                    file  \n 0  cd708819d3f5103afd6460b15e74eaf3.npz  }"},"metadata":{}}]},{"cell_type":"markdown","source":"# 📦 Define Dataset and Model","metadata":{"papermill":{"duration":0.008874,"end_time":"2023-09-01T16:56:21.968592","exception":false,"start_time":"2023-09-01T16:56:21.959718","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# 📦 Define Custom Dataset Class\n# This class, 'TileDataset', is a custom dataset class for our machine learning task.\n# It inherits from the PyTorch 'Dataset' class and implements the necessary methods (__init__, __len__, and __getitem__).\n\nclass TileDataset(Dataset):\n    def __init__(self, df):\n        self.df = df  # 💼 Initialize the dataset with a DataFrame containing the data\n\n    def __len__(self):\n        return len(self.df)  # 🔢 Define the length of the dataset, which is the number of rows in the DataFrame\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]  # 📄 Get a specific row from the DataFrame based on the provided index\n        config_feat = torch.tensor(row['config_feat'].astype(np.float32))  # 🧮 Convert and store 'config_feat' as a PyTorch tensor\n        node_feat = torch.tensor(row['node_feat'].astype(np.float32))  # 🧮 Convert and store 'node_feat' as a PyTorch tensor\n        node_opcode = torch.tensor(row['node_opcode'].astype(np.int32))  # 🧮 Convert and store 'node_opcode' as a PyTorch tensor\n        edge_index = torch.tensor(np.swapaxes(row['edge_index'],0,1).astype(np.int32))  # 🧮 Convert and store 'edge_index' as a PyTorch tensor with axis swapping\n        target = (row['config_runtime'] / (row['config_runtime_normalizers'] + 1e-5)).astype(np.float32)  # 📈 Calculate and store the target value with preprocessing\n        # 📊 Min-max scale the target value to ensure it's within a specific range (standardization)\n        target = (target - np.mean(target)) / (np.std(target) + 1e-5)\n        target = torch.tensor(target)  # 🧮 Convert and store the target as a PyTorch tensor\n        return config_feat, node_feat, node_opcode, edge_index, target  # 🔁 Return the data and target for a specific sample\n\n# This class defines the structure of our custom dataset, converting and preprocessing data as necessary for training and evaluation.\n# The relevant emojis provide a visual context for each part of the code.\n","metadata":{"papermill":{"duration":0.020329,"end_time":"2023-09-01T16:56:21.997734","exception":false,"start_time":"2023-09-01T16:56:21.977405","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-09-13T21:09:03.161828Z","iopub.execute_input":"2023-09-13T21:09:03.162613Z","iopub.status.idle":"2023-09-13T21:09:03.173596Z","shell.execute_reply.started":"2023-09-13T21:09:03.162575Z","shell.execute_reply":"2023-09-13T21:09:03.172164Z"},"trusted":true},"execution_count":5,"outputs":[]},{"cell_type":"markdown","source":"# 📦 Visualize some example tile data","metadata":{"execution":{"iopub.status.busy":"2023-09-13T17:45:08.506959Z","iopub.execute_input":"2023-09-13T17:45:08.507649Z","iopub.status.idle":"2023-09-13T17:45:08.512436Z","shell.execute_reply.started":"2023-09-13T17:45:08.507615Z","shell.execute_reply":"2023-09-13T17:45:08.511382Z"}}},{"cell_type":"code","source":"instruction_to_opcode = {\n    1: \"abs\",\n    2: \"add\",\n    3: \"add-dependency\",\n    4: \"after-all\",\n    5: \"all-reduce\",\n    6: \"all-to-all\",\n    7: \"atan2\",\n    8: \"batch-norm-grad\",\n    9: \"batch-norm-inference\",\n    10: \"batch-norm-training\",\n    11: \"bitcast\",\n    12: \"bitcast-convert\",\n    13: \"broadcast\",\n    14: \"call\",\n    15: \"ceil\",\n    16: \"cholesky\",\n    17: \"clamp\",\n    18: \"collective-permute\",\n    19: \"count-leading-zeros\",\n    20: \"compare\",\n    21: \"complex\",\n    22: \"concatenate\",\n    23: \"conditional\",\n    24: \"constant\",\n    25: \"convert\",\n    26: \"convolution\",\n    27: \"copy\",\n    28: \"copy-done\",\n    29: \"copy-start\",\n    30: \"cosine\",\n    31: \"custom-call\",\n    32: \"divide\",\n    33: \"domain\",\n    34: \"dot\",\n    35: \"dynamic-slice\",\n    36: \"dynamic-update-slice\",\n    37: \"exponential\",\n    38: \"exponential-minus-one\",\n    39: \"fft\",\n    40: \"floor\",\n    41: \"fusion\",\n    42: \"gather\",\n    43: \"get-dimension-size\",\n    44: \"set-dimension-size\",\n    45: \"get-tuple-element\",\n    46: \"imag\",\n    47: \"infeed\",\n    48: \"iota\",\n    49: \"is-finite\",\n    50: \"log\",\n    51: \"log-plus-one\",\n    52: \"and\",\n    53: \"not\",\n    54: \"or\",\n    55: \"xor\",\n    56: \"map\",\n    57: \"maximum\",\n    58: \"minimum\",\n    59: \"multiply\",\n    60: \"negate\",\n    61: \"outfeed\",\n    62: \"pad\",\n    63: \"parameter\",\n    64: \"partition-id\",\n    65: \"popcnt\",\n    66: \"power\",\n    67: \"real\",\n    68: \"recv\",\n    69: \"recv-done\",\n    70: \"reduce\",\n    71: \"reduce-precision\",\n    72: \"reduce-window\",\n    73: \"remainder\",\n    74: \"replica-id\",\n    75: \"reshape\",\n    76: \"reverse\",\n    77: \"rng\",\n    78: \"rng-get-and-update-state\",\n    79: \"rng-bit-generator\",\n    80: \"round-nearest-afz\",\n    81: \"rsqrt\",\n    82: \"scatter\",\n    83: \"select\",\n    84: \"select-and-scatter\",\n    85: \"send\",\n    86: \"send-done\",\n    87: \"shift-left\",\n    88: \"shift-right-arithmetic\",\n    89: \"shift-right-logical\",\n    90: \"sign\",\n    91: \"sine\",\n    92: \"slice\",\n    93: \"sort\",\n    94: \"sqrt\",\n    95: \"subtract\",\n    96: \"tanh\",\n    98: \"transpose\",\n    99: \"triangular-solve\",\n    100: \"tuple\",\n    102: \"while\",\n    103: \"cbrt\",\n    104: \"all-gather\",\n    105: \"collective-permute-start\",\n    106: \"collective-permute-done\",\n    107: \"logistic\",\n    108: \"dynamic-reshape\",\n    109: \"all-reduce-start\",\n    110: \"all-reduce-done\",\n    111: \"reduce-scatter\",\n    112: \"all-gather-start\",\n    113: \"all-gather-done\",\n    114: \"opt-barrier\",\n    115: \"async-start\",\n    116: \"async-update\",\n    117: \"async-done\",\n    118: \"round-nearest-even\",\n    119: \"stochastic-convert\",\n    120: \"tan\"\n}\n","metadata":{"execution":{"iopub.status.busy":"2023-09-13T21:09:03.845616Z","iopub.execute_input":"2023-09-13T21:09:03.846037Z","iopub.status.idle":"2023-09-13T21:09:03.865247Z","shell.execute_reply.started":"2023-09-13T21:09:03.846005Z","shell.execute_reply":"2023-09-13T21:09:03.863866Z"},"trusted":true},"execution_count":6,"outputs":[]},{"cell_type":"code","source":"import graphviz\nfrom IPython.display import display\n\ndef visualize_graph(data, engine='dot', node_size='0.15', font_size='8', ranksep='0.3'):\n    dot = graphviz.Digraph(comment='The Round Table', engine=engine)\n    opcode_to_name = instruction_to_opcode  # Assume this is the dictionary we constructed earlier\n\n    # Graph attributes to adjust spacing\n    dot.attr(ranksep=ranksep)   # Adjust as needed\n    dot.attr(nodesep='0.25')  # Adjust as needed\n\n    # Add nodes\n    for i, opcode in enumerate(data['node_opcode']):\n        instruction_name = opcode_to_name[opcode]\n        \n        # show shape\n        node_feat = data['node_feat'][i]\n        end_index = 26\n        while end_index > 21:\n            if node_feat[end_index] != 0:\n                break;\n            end_index -= 1\n        \n        dot.node(str(i), f\"{i}: {instruction_name}\\n{node_feat[21:end_index]}\", shape=\"box\", fontsize=font_size, height=node_size, width=node_size)\n\n    # Add edges\n    for u, v in data['edge_index']:\n        dot.edge(str(v), str(u))  # Note: the edge direction is reversed as per your requirement\n\n    # Display in Jupyter Notebook\n    display(dot)\n\n# # Sample data\n# data = {\n#     \"node_opcode\": [1, 2, 3],\n#     \"edge_index\": [[0, 1], [1, 2]]\n# }\n\n# visualize_graph(data, engine='dot', node_size='0.2', font_size='8')\n","metadata":{"execution":{"iopub.status.busy":"2023-09-13T21:09:04.494045Z","iopub.execute_input":"2023-09-13T21:09:04.495408Z","iopub.status.idle":"2023-09-13T21:09:04.548541Z","shell.execute_reply.started":"2023-09-13T21:09:04.495358Z","shell.execute_reply":"2023-09-13T21:09:04.547472Z"},"trusted":true},"execution_count":7,"outputs":[]},{"cell_type":"code","source":"tile_xla_for_vis = tile_xla['train']\nvisualize_id = 0\n\nprint(f\"Number of tile train data points: {len(tile_xla_for_vis)}\\n\")\n\n# an example data point:\ntile_xla_for_vis.iloc[visualize_id]","metadata":{"execution":{"iopub.status.busy":"2023-09-13T21:01:26.090905Z","iopub.execute_input":"2023-09-13T21:01:26.09178Z","iopub.status.idle":"2023-09-13T21:01:26.164064Z","shell.execute_reply.started":"2023-09-13T21:01:26.091735Z","shell.execute_reply":"2023-09-13T21:01:26.162693Z"},"trusted":true},"execution_count":12,"outputs":[{"name":"stdout","text":"Number of tile train data points: 5709\n\n","output_type":"stream"},{"execution_count":12,"output_type":"execute_result","data":{"text/plain":"node_feat                     [[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0,...\nnode_opcode                   [63, 63, 2, 63, 63, 2, 63, 11, 63, 11, 63, 24,...\nedge_index                    [[2, 0], [2, 1], [5, 3], [5, 4], [7, 6], [9, 8...\nconfig_feat                   [[16.0, 16.0, 1.0, 2.0, 0.0, 0.0, 35.0, 512.0,...\nconfig_runtime                [238408, 5052251, 1873104, 2452158, 1872430, 1...\nconfig_runtime_normalizers    [238408, 238408, 238408, 238408, 238408, 23840...\nfile                                   retinanet.4x4.fp32_-431a58cc30e72ec6.npz\nName: 0, dtype: object"},"metadata":{}}]},{"cell_type":"code","source":"visualize_graph(tile_xla_for_vis.iloc[visualize_id])","metadata":{"execution":{"iopub.status.busy":"2023-09-13T21:01:26.46269Z","iopub.execute_input":"2023-09-13T21:01:26.463963Z","iopub.status.idle":"2023-09-13T21:01:26.512277Z","shell.execute_reply.started":"2023-09-13T21:01:26.463912Z","shell.execute_reply":"2023-09-13T21:01:26.51028Z"},"trusted":true},"execution_count":13,"outputs":[{"output_type":"display_data","data":{"image/svg+xml":"<?xml version=\"1.0\" encoding=\"UTF-8\" standalone=\"no\"?>\n<!DOCTYPE svg PUBLIC \"-//W3C//DTD SVG 1.1//EN\"\n \"http://www.w3.org/Graphics/SVG/1.1/DTD/svg11.dtd\">\n<!-- Generated by graphviz version 8.0.5 (20230506.1012)\n -->\n<!-- Pages: 1 -->\n<svg width=\"611pt\" height=\"481pt\"\n viewBox=\"0.00 0.00 611.38 481.00\" xmlns=\"http://www.w3.org/2000/svg\" xmlns:xlink=\"http://www.w3.org/1999/xlink\">\n<g id=\"graph0\" class=\"graph\" transform=\"scale(1 1) rotate(0) translate(4 477)\">\n<polygon fill=\"white\" stroke=\"none\" points=\"-4,4 -4,-477 607.38,-477 607.38,4 -4,4\"/>\n<!-- 0 -->\n<g id=\"node1\" class=\"node\">\n<title>0</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"172.75,-225.5 105,-225.5 105,-198 172.75,-198 172.75,-225.5\"/>\n<text text-anchor=\"middle\" x=\"138.88\" y=\"-213.9\" font-family=\"Times,serif\" font-size=\"8.00\">0: parameter</text>\n<text text-anchor=\"middle\" x=\"138.88\" y=\"-204.15\" font-family=\"Times,serif\" font-size=\"8.00\">[]</text>\n</g>\n<!-- 2 -->\n<g id=\"node3\" class=\"node\">\n<title>2</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"245.62,-176 204.12,-176 204.12,-148.5 245.62,-148.5 245.62,-176\"/>\n<text text-anchor=\"middle\" x=\"224.88\" y=\"-164.4\" font-family=\"Times,serif\" font-size=\"8.00\">2: add</text>\n<text text-anchor=\"middle\" x=\"224.88\" y=\"-154.65\" font-family=\"Times,serif\" font-size=\"8.00\">[]</text>\n</g>\n<!-- 0&#45;&gt;2 -->\n<g id=\"edge1\" class=\"edge\">\n<title>0&#45;&gt;2</title>\n<path fill=\"none\" stroke=\"black\" d=\"M162.83,-197.52C172.55,-192.15 183.89,-185.89 194.2,-180.19\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"195.62,-182.85 202.68,-174.95 192.24,-176.73 195.62,-182.85\"/>\n</g>\n<!-- 1 -->\n<g id=\"node2\" class=\"node\">\n<title>1</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"258.75,-225.5 191,-225.5 191,-198 258.75,-198 258.75,-225.5\"/>\n<text text-anchor=\"middle\" x=\"224.88\" y=\"-213.9\" font-family=\"Times,serif\" font-size=\"8.00\">1: parameter</text>\n<text text-anchor=\"middle\" x=\"224.88\" y=\"-204.15\" font-family=\"Times,serif\" font-size=\"8.00\">[]</text>\n</g>\n<!-- 1&#45;&gt;2 -->\n<g id=\"edge2\" class=\"edge\">\n<title>1&#45;&gt;2</title>\n<path fill=\"none\" stroke=\"black\" d=\"M224.88,-197.52C224.88,-194.33 224.88,-190.82 224.88,-187.3\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"228.38,-187.41 224.88,-177.41 221.38,-187.41 228.38,-187.41\"/>\n</g>\n<!-- 23 -->\n<g id=\"node24\" class=\"node\">\n<title>23</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"333.62,-126.5 274.12,-126.5 274.12,-99 333.62,-99 333.62,-126.5\"/>\n<text text-anchor=\"middle\" x=\"303.88\" y=\"-114.9\" font-family=\"Times,serif\" font-size=\"8.00\">23: reduce</text>\n<text text-anchor=\"middle\" x=\"303.88\" y=\"-105.15\" font-family=\"Times,serif\" font-size=\"8.00\">[]</text>\n</g>\n<!-- 2&#45;&gt;23 -->\n<g id=\"edge20\" class=\"edge\">\n<title>2&#45;&gt;23</title>\n<path fill=\"none\" stroke=\"black\" d=\"M246.05,-148.52C254.24,-143.6 263.78,-137.86 272.72,-132.48\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"274.33,-135 281.1,-126.85 270.72,-129 274.33,-135\"/>\n</g>\n<!-- 3 -->\n<g id=\"node4\" class=\"node\">\n<title>3</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"430.75,-225.5 363,-225.5 363,-198 430.75,-198 430.75,-225.5\"/>\n<text text-anchor=\"middle\" x=\"396.88\" y=\"-213.9\" font-family=\"Times,serif\" font-size=\"8.00\">3: parameter</text>\n<text text-anchor=\"middle\" x=\"396.88\" y=\"-204.15\" font-family=\"Times,serif\" font-size=\"8.00\">[]</text>\n</g>\n<!-- 5 -->\n<g id=\"node6\" class=\"node\">\n<title>5</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"481.62,-176 440.12,-176 440.12,-148.5 481.62,-148.5 481.62,-176\"/>\n<text text-anchor=\"middle\" x=\"460.88\" y=\"-164.4\" font-family=\"Times,serif\" font-size=\"8.00\">5: add</text>\n<text text-anchor=\"middle\" x=\"460.88\" y=\"-154.65\" font-family=\"Times,serif\" font-size=\"8.00\">[]</text>\n</g>\n<!-- 3&#45;&gt;5 -->\n<g id=\"edge3\" class=\"edge\">\n<title>3&#45;&gt;5</title>\n<path fill=\"none\" stroke=\"black\" d=\"M414.7,-197.52C420.81,-192.99 427.77,-187.82 434.41,-182.89\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"436.2,-185.18 442.14,-176.41 432.03,-179.56 436.2,-185.18\"/>\n</g>\n<!-- 4 -->\n<g id=\"node5\" class=\"node\">\n<title>4</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"516.75,-225.5 449,-225.5 449,-198 516.75,-198 516.75,-225.5\"/>\n<text text-anchor=\"middle\" x=\"482.88\" y=\"-213.9\" font-family=\"Times,serif\" font-size=\"8.00\">4: parameter</text>\n<text text-anchor=\"middle\" x=\"482.88\" y=\"-204.15\" font-family=\"Times,serif\" font-size=\"8.00\">[]</text>\n</g>\n<!-- 4&#45;&gt;5 -->\n<g id=\"edge4\" class=\"edge\">\n<title>4&#45;&gt;5</title>\n<path fill=\"none\" stroke=\"black\" d=\"M476.75,-197.52C475.13,-194.02 473.33,-190.14 471.54,-186.28\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"474.35,-185.01 466.97,-177.41 468,-187.96 474.35,-185.01\"/>\n</g>\n<!-- 25 -->\n<g id=\"node26\" class=\"node\">\n<title>25</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"454.62,-126.5 395.12,-126.5 395.12,-99 454.62,-99 454.62,-126.5\"/>\n<text text-anchor=\"middle\" x=\"424.88\" y=\"-114.9\" font-family=\"Times,serif\" font-size=\"8.00\">25: reduce</text>\n<text text-anchor=\"middle\" x=\"424.88\" y=\"-105.15\" font-family=\"Times,serif\" font-size=\"8.00\">[]</text>\n</g>\n<!-- 5&#45;&gt;25 -->\n<g id=\"edge24\" class=\"edge\">\n<title>5&#45;&gt;25</title>\n<path fill=\"none\" stroke=\"black\" d=\"M450.85,-148.02C447.88,-144.11 444.56,-139.72 441.3,-135.42\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"443.68,-133.77 434.85,-127.91 438.1,-137.99 443.68,-133.77\"/>\n</g>\n<!-- 6 -->\n<g id=\"node7\" class=\"node\">\n<title>6</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"67.75,-473 0,-473 0,-445.5 67.75,-445.5 67.75,-473\"/>\n<text text-anchor=\"middle\" x=\"33.88\" y=\"-461.4\" font-family=\"Times,serif\" font-size=\"8.00\">6: parameter</text>\n<text text-anchor=\"middle\" x=\"33.88\" y=\"-451.65\" font-family=\"Times,serif\" font-size=\"8.00\">[ 8. 80. 80.]</text>\n</g>\n<!-- 7 -->\n<g id=\"node8\" class=\"node\">\n<title>7</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"64.75,-423.5 3,-423.5 3,-396 64.75,-396 64.75,-423.5\"/>\n<text text-anchor=\"middle\" x=\"33.88\" y=\"-411.9\" font-family=\"Times,serif\" font-size=\"8.00\">7: bitcast</text>\n<text text-anchor=\"middle\" x=\"33.88\" y=\"-402.15\" font-family=\"Times,serif\" font-size=\"8.00\">[ 8. 80. 80.]</text>\n</g>\n<!-- 6&#45;&gt;7 -->\n<g id=\"edge5\" class=\"edge\">\n<title>6&#45;&gt;7</title>\n<path fill=\"none\" stroke=\"black\" d=\"M33.88,-445.02C33.88,-441.83 33.88,-438.32 33.88,-434.8\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"37.38,-434.91 33.88,-424.91 30.38,-434.91 37.38,-434.91\"/>\n</g>\n<!-- 15 -->\n<g id=\"node16\" class=\"node\">\n<title>15</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"149.75,-374 88,-374 88,-346.5 149.75,-346.5 149.75,-374\"/>\n<text text-anchor=\"middle\" x=\"118.88\" y=\"-362.4\" font-family=\"Times,serif\" font-size=\"8.00\">15: fusion</text>\n<text text-anchor=\"middle\" x=\"118.88\" y=\"-352.65\" font-family=\"Times,serif\" font-size=\"8.00\">[ 8. 80. 80.]</text>\n</g>\n<!-- 7&#45;&gt;15 -->\n<g id=\"edge10\" class=\"edge\">\n<title>7&#45;&gt;15</title>\n<path fill=\"none\" stroke=\"black\" d=\"M57.55,-395.52C66.3,-390.63 76.39,-384.99 85.82,-379.72\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"87.3,-382.34 94.32,-374.41 83.89,-376.23 87.3,-382.34\"/>\n</g>\n<!-- 8 -->\n<g id=\"node9\" class=\"node\">\n<title>8</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"237.75,-473 170,-473 170,-445.5 237.75,-445.5 237.75,-473\"/>\n<text text-anchor=\"middle\" x=\"203.88\" y=\"-461.4\" font-family=\"Times,serif\" font-size=\"8.00\">8: parameter</text>\n<text text-anchor=\"middle\" x=\"203.88\" y=\"-451.65\" font-family=\"Times,serif\" font-size=\"8.00\">[ 8. 80. 80.]</text>\n</g>\n<!-- 9 -->\n<g id=\"node10\" class=\"node\">\n<title>9</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"234.75,-423.5 173,-423.5 173,-396 234.75,-396 234.75,-423.5\"/>\n<text text-anchor=\"middle\" x=\"203.88\" y=\"-411.9\" font-family=\"Times,serif\" font-size=\"8.00\">9: bitcast</text>\n<text text-anchor=\"middle\" x=\"203.88\" y=\"-402.15\" font-family=\"Times,serif\" font-size=\"8.00\">[ 8. 80. 80.]</text>\n</g>\n<!-- 8&#45;&gt;9 -->\n<g id=\"edge6\" class=\"edge\">\n<title>8&#45;&gt;9</title>\n<path fill=\"none\" stroke=\"black\" d=\"M203.88,-445.02C203.88,-441.83 203.88,-438.32 203.88,-434.8\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"207.38,-434.91 203.88,-424.91 200.38,-434.91 207.38,-434.91\"/>\n</g>\n<!-- 17 -->\n<g id=\"node18\" class=\"node\">\n<title>17</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"234.75,-374 173,-374 173,-346.5 234.75,-346.5 234.75,-374\"/>\n<text text-anchor=\"middle\" x=\"203.88\" y=\"-362.4\" font-family=\"Times,serif\" font-size=\"8.00\">17: fusion</text>\n<text text-anchor=\"middle\" x=\"203.88\" y=\"-352.65\" font-family=\"Times,serif\" font-size=\"8.00\">[ 8. 80. 80.]</text>\n</g>\n<!-- 9&#45;&gt;17 -->\n<g id=\"edge12\" class=\"edge\">\n<title>9&#45;&gt;17</title>\n<path fill=\"none\" stroke=\"black\" d=\"M203.88,-395.52C203.88,-392.33 203.88,-388.82 203.88,-385.3\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"207.38,-385.41 203.88,-375.41 200.38,-385.41 207.38,-385.41\"/>\n</g>\n<!-- 10 -->\n<g id=\"node11\" class=\"node\">\n<title>10</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"491.38,-324.5 418.38,-324.5 418.38,-297 491.38,-297 491.38,-324.5\"/>\n<text text-anchor=\"middle\" x=\"454.88\" y=\"-312.9\" font-family=\"Times,serif\" font-size=\"8.00\">10: parameter</text>\n<text text-anchor=\"middle\" x=\"454.88\" y=\"-303.15\" font-family=\"Times,serif\" font-size=\"8.00\">[ &#160;1. &#160;&#160;1. 512.]</text>\n</g>\n<!-- 13 -->\n<g id=\"node14\" class=\"node\">\n<title>13</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"363.12,-275 294.62,-275 294.62,-247.5 363.12,-247.5 363.12,-275\"/>\n<text text-anchor=\"middle\" x=\"328.88\" y=\"-263.4\" font-family=\"Times,serif\" font-size=\"8.00\">13: multiply</text>\n<text text-anchor=\"middle\" x=\"328.88\" y=\"-253.65\" font-family=\"Times,serif\" font-size=\"8.00\">[ &#160;1. &#160;&#160;1. 512.]</text>\n</g>\n<!-- 10&#45;&gt;13 -->\n<g id=\"edge8\" class=\"edge\">\n<title>10&#45;&gt;13</title>\n<path fill=\"none\" stroke=\"black\" d=\"M419.78,-296.52C405.46,-291.12 388.73,-284.82 373.56,-279.09\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"374.94,-275.5 364.35,-275.25 372.47,-282.05 374.94,-275.5\"/>\n</g>\n<!-- 24 -->\n<g id=\"node25\" class=\"node\">\n<title>24</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"561.12,-275 492.62,-275 492.62,-247.5 561.12,-247.5 561.12,-275\"/>\n<text text-anchor=\"middle\" x=\"526.88\" y=\"-263.4\" font-family=\"Times,serif\" font-size=\"8.00\">24: multiply</text>\n<text text-anchor=\"middle\" x=\"526.88\" y=\"-253.65\" font-family=\"Times,serif\" font-size=\"8.00\">[ &#160;1. &#160;&#160;1. 512.]</text>\n</g>\n<!-- 10&#45;&gt;24 -->\n<g id=\"edge23\" class=\"edge\">\n<title>10&#45;&gt;24</title>\n<path fill=\"none\" stroke=\"black\" d=\"M474.93,-296.52C482.03,-291.83 490.17,-286.47 497.86,-281.39\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"499.51,-283.84 505.93,-275.41 495.65,-278 499.51,-283.84\"/>\n</g>\n<!-- 11 -->\n<g id=\"node12\" class=\"node\">\n<title>11</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"361.62,-374 296.12,-374 296.12,-346.5 361.62,-346.5 361.62,-374\"/>\n<text text-anchor=\"middle\" x=\"328.88\" y=\"-362.4\" font-family=\"Times,serif\" font-size=\"8.00\">11: constant</text>\n<text text-anchor=\"middle\" x=\"328.88\" y=\"-352.65\" font-family=\"Times,serif\" font-size=\"8.00\">[]</text>\n</g>\n<!-- 12 -->\n<g id=\"node13\" class=\"node\">\n<title>12</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"364.25,-324.5 293.5,-324.5 293.5,-297 364.25,-297 364.25,-324.5\"/>\n<text text-anchor=\"middle\" x=\"328.88\" y=\"-312.9\" font-family=\"Times,serif\" font-size=\"8.00\">12: broadcast</text>\n<text text-anchor=\"middle\" x=\"328.88\" y=\"-303.15\" font-family=\"Times,serif\" font-size=\"8.00\">[ &#160;1. &#160;&#160;1. 512.]</text>\n</g>\n<!-- 11&#45;&gt;12 -->\n<g id=\"edge7\" class=\"edge\">\n<title>11&#45;&gt;12</title>\n<path fill=\"none\" stroke=\"black\" d=\"M328.88,-346.02C328.88,-342.83 328.88,-339.32 328.88,-335.8\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"332.38,-335.91 328.88,-325.91 325.38,-335.91 332.38,-335.91\"/>\n</g>\n<!-- 12&#45;&gt;13 -->\n<g id=\"edge9\" class=\"edge\">\n<title>12&#45;&gt;13</title>\n<path fill=\"none\" stroke=\"black\" d=\"M328.88,-296.52C328.88,-293.33 328.88,-289.82 328.88,-286.3\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"332.38,-286.41 328.88,-276.41 325.38,-286.41 332.38,-286.41\"/>\n</g>\n<!-- 20 -->\n<g id=\"node21\" class=\"node\">\n<title>20</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"345.12,-225.5 276.62,-225.5 276.62,-198 345.12,-198 345.12,-225.5\"/>\n<text text-anchor=\"middle\" x=\"310.88\" y=\"-213.9\" font-family=\"Times,serif\" font-size=\"8.00\">20: add</text>\n<text text-anchor=\"middle\" x=\"310.88\" y=\"-204.15\" font-family=\"Times,serif\" font-size=\"8.00\">[ &#160;1. &#160;&#160;1. 512.]</text>\n</g>\n<!-- 13&#45;&gt;20 -->\n<g id=\"edge17\" class=\"edge\">\n<title>13&#45;&gt;20</title>\n<path fill=\"none\" stroke=\"black\" d=\"M323.86,-247.02C322.57,-243.62 321.15,-239.86 319.73,-236.12\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"322.68,-235.02 315.86,-226.91 316.13,-237.5 322.68,-235.02\"/>\n</g>\n<!-- 14 -->\n<g id=\"node15\" class=\"node\">\n<title>14</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"155.38,-423.5 82.38,-423.5 82.38,-396 155.38,-396 155.38,-423.5\"/>\n<text text-anchor=\"middle\" x=\"118.88\" y=\"-411.9\" font-family=\"Times,serif\" font-size=\"8.00\">14: parameter</text>\n<text text-anchor=\"middle\" x=\"118.88\" y=\"-402.15\" font-family=\"Times,serif\" font-size=\"8.00\">[ 8. 80. 80.]</text>\n</g>\n<!-- 14&#45;&gt;15 -->\n<g id=\"edge11\" class=\"edge\">\n<title>14&#45;&gt;15</title>\n<path fill=\"none\" stroke=\"black\" d=\"M118.88,-395.52C118.88,-392.33 118.88,-388.82 118.88,-385.3\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"122.38,-385.41 118.88,-375.41 115.38,-385.41 122.38,-385.41\"/>\n</g>\n<!-- 18 -->\n<g id=\"node19\" class=\"node\">\n<title>18</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"177.62,-324.5 100.12,-324.5 100.12,-297 177.62,-297 177.62,-324.5\"/>\n<text text-anchor=\"middle\" x=\"138.88\" y=\"-312.9\" font-family=\"Times,serif\" font-size=\"8.00\">18: convolution</text>\n<text text-anchor=\"middle\" x=\"138.88\" y=\"-303.15\" font-family=\"Times,serif\" font-size=\"8.00\">[ &#160;1. &#160;&#160;1. 512.]</text>\n</g>\n<!-- 15&#45;&gt;18 -->\n<g id=\"edge14\" class=\"edge\">\n<title>15&#45;&gt;18</title>\n<path fill=\"none\" stroke=\"black\" d=\"M124.45,-346.02C125.92,-342.52 127.55,-338.64 129.18,-334.78\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"132.68,-336.49 133.33,-325.91 126.23,-333.77 132.68,-336.49\"/>\n</g>\n<!-- 16 -->\n<g id=\"node17\" class=\"node\">\n<title>16</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"325.38,-423.5 252.38,-423.5 252.38,-396 325.38,-396 325.38,-423.5\"/>\n<text text-anchor=\"middle\" x=\"288.88\" y=\"-411.9\" font-family=\"Times,serif\" font-size=\"8.00\">16: parameter</text>\n<text text-anchor=\"middle\" x=\"288.88\" y=\"-402.15\" font-family=\"Times,serif\" font-size=\"8.00\">[ 8. 80. 80.]</text>\n</g>\n<!-- 16&#45;&gt;17 -->\n<g id=\"edge13\" class=\"edge\">\n<title>16&#45;&gt;17</title>\n<path fill=\"none\" stroke=\"black\" d=\"M265.2,-395.52C256.45,-390.63 246.36,-384.99 236.93,-379.72\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"238.86,-376.23 228.43,-374.41 235.45,-382.34 238.86,-376.23\"/>\n</g>\n<!-- 17&#45;&gt;18 -->\n<g id=\"edge15\" class=\"edge\">\n<title>17&#45;&gt;18</title>\n<path fill=\"none\" stroke=\"black\" d=\"M185.77,-346.02C179.57,-341.49 172.5,-336.32 165.76,-331.39\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"168.02,-327.98 157.88,-324.91 163.89,-333.64 168.02,-327.98\"/>\n</g>\n<!-- 19 -->\n<g id=\"node20\" class=\"node\">\n<title>19</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"236.12,-275 167.62,-275 167.62,-247.5 236.12,-247.5 236.12,-275\"/>\n<text text-anchor=\"middle\" x=\"201.88\" y=\"-263.4\" font-family=\"Times,serif\" font-size=\"8.00\">19: convert</text>\n<text text-anchor=\"middle\" x=\"201.88\" y=\"-253.65\" font-family=\"Times,serif\" font-size=\"8.00\">[ &#160;1. &#160;&#160;1. 512.]</text>\n</g>\n<!-- 18&#45;&gt;19 -->\n<g id=\"edge16\" class=\"edge\">\n<title>18&#45;&gt;19</title>\n<path fill=\"none\" stroke=\"black\" d=\"M156.42,-296.52C162.43,-291.99 169.29,-286.82 175.82,-281.89\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"177.54,-284.23 183.42,-275.41 173.33,-278.64 177.54,-284.23\"/>\n</g>\n<!-- 26 -->\n<g id=\"node27\" class=\"node\">\n<title>26</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"329.88,-77 277.88,-77 277.88,-49.5 329.88,-49.5 329.88,-77\"/>\n<text text-anchor=\"middle\" x=\"303.88\" y=\"-65.4\" font-family=\"Times,serif\" font-size=\"8.00\">26: tuple</text>\n<text text-anchor=\"middle\" x=\"303.88\" y=\"-55.65\" font-family=\"Times,serif\" font-size=\"8.00\">[]</text>\n</g>\n<!-- 18&#45;&gt;26 -->\n<g id=\"edge27\" class=\"edge\">\n<title>18&#45;&gt;26</title>\n<path fill=\"none\" stroke=\"black\" d=\"M123.55,-296.68C105.19,-279.47 76.88,-247.33 76.88,-212.75 76.88,-212.75 76.88,-212.75 76.88,-161.25 76.88,-79.78 202.03,-65.82 266.48,-64.03\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"266.46,-67.51 276.39,-63.83 266.33,-60.51 266.46,-67.51\"/>\n</g>\n<!-- 19&#45;&gt;20 -->\n<g id=\"edge18\" class=\"edge\">\n<title>19&#45;&gt;20</title>\n<path fill=\"none\" stroke=\"black\" d=\"M232.24,-247.02C244.17,-241.82 258.03,-235.78 270.76,-230.23\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"271.91,-233.11 279.68,-225.91 269.11,-226.7 271.91,-233.11\"/>\n</g>\n<!-- 21 -->\n<g id=\"node22\" class=\"node\">\n<title>21</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"338.12,-176 269.62,-176 269.62,-148.5 338.12,-148.5 338.12,-176\"/>\n<text text-anchor=\"middle\" x=\"303.88\" y=\"-164.4\" font-family=\"Times,serif\" font-size=\"8.00\">21: multiply</text>\n<text text-anchor=\"middle\" x=\"303.88\" y=\"-154.65\" font-family=\"Times,serif\" font-size=\"8.00\">[ &#160;1. &#160;&#160;1. 512.]</text>\n</g>\n<!-- 20&#45;&gt;21 -->\n<g id=\"edge19\" class=\"edge\">\n<title>20&#45;&gt;21</title>\n<path fill=\"none\" stroke=\"black\" d=\"M308.93,-197.52C308.44,-194.22 307.9,-190.59 307.37,-186.96\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"310.74,-186.79 305.81,-177.41 303.81,-187.81 310.74,-186.79\"/>\n</g>\n<!-- 21&#45;&gt;23 -->\n<g id=\"edge21\" class=\"edge\">\n<title>21&#45;&gt;23</title>\n<path fill=\"none\" stroke=\"black\" d=\"M303.88,-148.02C303.88,-144.83 303.88,-141.32 303.88,-137.8\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"307.38,-137.91 303.88,-127.91 300.38,-137.91 307.38,-137.91\"/>\n</g>\n<!-- 22 -->\n<g id=\"node23\" class=\"node\">\n<title>22</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"421.62,-176 356.12,-176 356.12,-148.5 421.62,-148.5 421.62,-176\"/>\n<text text-anchor=\"middle\" x=\"388.88\" y=\"-164.4\" font-family=\"Times,serif\" font-size=\"8.00\">22: constant</text>\n<text text-anchor=\"middle\" x=\"388.88\" y=\"-154.65\" font-family=\"Times,serif\" font-size=\"8.00\">[]</text>\n</g>\n<!-- 22&#45;&gt;23 -->\n<g id=\"edge22\" class=\"edge\">\n<title>22&#45;&gt;23</title>\n<path fill=\"none\" stroke=\"black\" d=\"M365.2,-148.02C356.45,-143.13 346.36,-137.49 336.93,-132.22\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"338.86,-128.73 328.43,-126.91 335.45,-134.84 338.86,-128.73\"/>\n</g>\n<!-- 22&#45;&gt;25 -->\n<g id=\"edge25\" class=\"edge\">\n<title>22&#45;&gt;25</title>\n<path fill=\"none\" stroke=\"black\" d=\"M398.9,-148.02C401.87,-144.11 405.19,-139.72 408.45,-135.42\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"411.65,-137.99 414.9,-127.91 406.07,-133.77 411.65,-137.99\"/>\n</g>\n<!-- 23&#45;&gt;26 -->\n<g id=\"edge28\" class=\"edge\">\n<title>23&#45;&gt;26</title>\n<path fill=\"none\" stroke=\"black\" d=\"M303.88,-98.52C303.88,-95.33 303.88,-91.82 303.88,-88.3\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"307.38,-88.41 303.88,-78.41 300.38,-88.41 307.38,-88.41\"/>\n</g>\n<!-- 24&#45;&gt;25 -->\n<g id=\"edge26\" class=\"edge\">\n<title>24&#45;&gt;25</title>\n<path fill=\"none\" stroke=\"black\" d=\"M528.88,-247C530.32,-234.14 531.18,-214.23 525.88,-198 517.51,-172.39 511.13,-166.27 490.88,-148.5 483.15,-141.73 473.88,-135.82 464.76,-130.88\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"466.5,-127.33 456,-125.92 463.33,-133.57 466.5,-127.33\"/>\n</g>\n<!-- 25&#45;&gt;26 -->\n<g id=\"edge29\" class=\"edge\">\n<title>25&#45;&gt;26</title>\n<path fill=\"none\" stroke=\"black\" d=\"M394.65,-99.89C378.04,-93.36 357.31,-85.23 339.86,-78.38\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"341.61,-74.91 331.02,-74.51 339.05,-81.43 341.61,-74.91\"/>\n</g>\n<!-- 30 -->\n<g id=\"node31\" class=\"node\">\n<title>30</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"457.75,-27.5 402,-27.5 402,0 457.75,0 457.75,-27.5\"/>\n<text text-anchor=\"middle\" x=\"429.88\" y=\"-15.9\" font-family=\"Times,serif\" font-size=\"8.00\">30: fusion</text>\n<text text-anchor=\"middle\" x=\"429.88\" y=\"-6.15\" font-family=\"Times,serif\" font-size=\"8.00\">[]</text>\n</g>\n<!-- 26&#45;&gt;30 -->\n<g id=\"edge30\" class=\"edge\">\n<title>26&#45;&gt;30</title>\n<path fill=\"none\" stroke=\"black\" d=\"M330.29,-52.29C348.03,-45.61 371.69,-36.69 391.5,-29.22\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"392.62,-32.16 400.75,-25.35 390.15,-25.61 392.62,-32.16\"/>\n</g>\n<!-- 27 -->\n<g id=\"node28\" class=\"node\">\n<title>27</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"421.38,-77 348.38,-77 348.38,-49.5 421.38,-49.5 421.38,-77\"/>\n<text text-anchor=\"middle\" x=\"384.88\" y=\"-65.4\" font-family=\"Times,serif\" font-size=\"8.00\">27: parameter</text>\n<text text-anchor=\"middle\" x=\"384.88\" y=\"-55.65\" font-family=\"Times,serif\" font-size=\"8.00\">[ &#160;1. &#160;&#160;1. 512.]</text>\n</g>\n<!-- 27&#45;&gt;30 -->\n<g id=\"edge31\" class=\"edge\">\n<title>27&#45;&gt;30</title>\n<path fill=\"none\" stroke=\"black\" d=\"M397.41,-49.02C401.31,-44.9 405.71,-40.25 409.99,-35.74\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"412.07,-38.58 416.41,-28.91 406.99,-33.76 412.07,-38.58\"/>\n</g>\n<!-- 28 -->\n<g id=\"node29\" class=\"node\">\n<title>28</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"512.38,-77 439.38,-77 439.38,-49.5 512.38,-49.5 512.38,-77\"/>\n<text text-anchor=\"middle\" x=\"475.88\" y=\"-65.4\" font-family=\"Times,serif\" font-size=\"8.00\">28: parameter</text>\n<text text-anchor=\"middle\" x=\"475.88\" y=\"-55.65\" font-family=\"Times,serif\" font-size=\"8.00\">[ 8. 80. 80.]</text>\n</g>\n<!-- 28&#45;&gt;30 -->\n<g id=\"edge32\" class=\"edge\">\n<title>28&#45;&gt;30</title>\n<path fill=\"none\" stroke=\"black\" d=\"M463.06,-49.02C459.07,-44.9 454.57,-40.25 450.21,-35.74\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"453.09,-33.66 443.62,-28.91 448.06,-38.53 453.09,-33.66\"/>\n</g>\n<!-- 29 -->\n<g id=\"node30\" class=\"node\">\n<title>29</title>\n<polygon fill=\"none\" stroke=\"black\" points=\"603.38,-77 530.38,-77 530.38,-49.5 603.38,-49.5 603.38,-77\"/>\n<text text-anchor=\"middle\" x=\"566.88\" y=\"-65.4\" font-family=\"Times,serif\" font-size=\"8.00\">29: parameter</text>\n<text text-anchor=\"middle\" x=\"566.88\" y=\"-55.65\" font-family=\"Times,serif\" font-size=\"8.00\">[ 8. 80. 80.]</text>\n</g>\n<!-- 29&#45;&gt;30 -->\n<g id=\"edge33\" class=\"edge\">\n<title>29&#45;&gt;30</title>\n<path fill=\"none\" stroke=\"black\" d=\"M530.16,-49.52C511.04,-42.89 487.66,-34.79 468.29,-28.07\"/>\n<polygon fill=\"black\" stroke=\"black\" points=\"469.59,-24.47 458.99,-24.5 467.3,-31.08 469.59,-24.47\"/>\n</g>\n</g>\n</svg>\n","text/plain":"<graphviz.graphs.Digraph at 0x780b4c098b20>"},"metadata":{}}]},{"cell_type":"code","source":"layout_xla_for_vis = layout_xla_default['train']\nvisualize_id = 0\n\nprint(f\"Number of layout train data points: {len(layout_xla_for_vis)}\\n\")\n\n# an example data point:\nlayout_xla_for_vis.iloc[visualize_id]","metadata":{"execution":{"iopub.status.busy":"2023-09-13T21:09:12.459211Z","iopub.execute_input":"2023-09-13T21:09:12.459628Z","iopub.status.idle":"2023-09-13T21:09:15.206172Z","shell.execute_reply.started":"2023-09-13T21:09:12.459598Z","shell.execute_reply":"2023-09-13T21:09:15.205022Z"},"trusted":true},"execution_count":8,"outputs":[{"name":"stdout","text":"Number of layout train data points: 1\n\n","output_type":"stream"},{"execution_count":8,"output_type":"execute_result","data":{"text/plain":"node_feat           [[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0,...\nnode_opcode         [63, 63, 2, 63, 63, 57, 63, 63, 2, 63, 63, 2, ...\nedge_index          [[2, 0], [2, 1], [5, 3], [5, 4], [8, 6], [8, 7...\nnode_config_feat    [[[3.0, 0.0, -1.0, -1.0, -1.0, -1.0, 0.0, 3.0,...\nnode_config_ids     [1001, 1034, 1062, 1087, 1113, 1140, 1165, 119...\nconfig_runtime      [73585984, 75217964, 73773521, 73779258, 73778...\nnode_splits                                               [[0, 5605]]\nfile                                            resnet50.8x8.fp16.npz\nName: 0, dtype: object"},"metadata":{}}]},{"cell_type":"code","source":"len(layout_xla_for_vis.iloc[visualize_id][\"node_feat\"])","metadata":{"execution":{"iopub.status.busy":"2023-09-13T21:10:09.954774Z","iopub.execute_input":"2023-09-13T21:10:09.956025Z","iopub.status.idle":"2023-09-13T21:10:09.964936Z","shell.execute_reply.started":"2023-09-13T21:10:09.955959Z","shell.execute_reply":"2023-09-13T21:10:09.963524Z"},"trusted":true},"execution_count":10,"outputs":[{"execution_count":10,"output_type":"execute_result","data":{"text/plain":"5605"},"metadata":{}}]},{"cell_type":"code","source":"visualize_graph(layout_xla_default[\"train\"].iloc[visualize_id])","metadata":{"execution":{"iopub.status.busy":"2023-09-13T21:10:21.378263Z","iopub.execute_input":"2023-09-13T21:10:21.378732Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 🧠 Define Simple Neural Network Model\n# In this cell, we define a simple neural network model named 'SimpleModel'.\n# This model takes input data with specified dimensions and passes it through convolutional and dense layers.\n\nclass SimpleModel(torch.nn.Module):\n    def __init__(self, hidden_channels, graph_feats, hidden_dim):\n        super().__init__()  # 🧬 Initialize the parent class 'torch.nn.Module'\n        \n        op_embedding_dim = 4  # I choose 4-dimensional embedding\n        self.embedding = torch.nn.Embedding(120,  # 120 different op-codes\n                                            op_embedding_dim,\n                                           )\n        assert len(hidden_channels) > 0\n        in_channels = op_embedding_dim + 140\n        self.convs = torch.nn.ModuleList()\n        last_dim = hidden_channels[0]\n        \n        # Create a sequence of Graph Convolutional Network (GCN) layers\n        self.convs.append(GCNConv(in_channels, hidden_channels[0]))\n        for i in range(len(hidden_channels) - 1):\n            self.convs.append(GCNConv(hidden_channels[i], hidden_channels[i+1]))\n            last_dim = hidden_channels[i+1]\n        self.convs.append(GCNConv(last_dim, graph_feats))\n        \n        # Define a sequential dense neural network\n        self.dense = torch.nn.Sequential(nn.Linear(graph_feats + 24, 64),\n                                         nn.ReLU(),\n                                         nn.Linear(64, 64),\n                                         nn.ReLU(),\n                                         nn.Linear(64, 1),\n                                        )\n\n    def forward(self, x_cfg: Tensor, x_feat: Tensor, x_op: Tensor, edge_index: Tensor) -> Tensor:\n        \n        # Get graph features\n        x = torch.cat([x_feat, self.embedding(x_op)], dim=1)  # 📊 Concatenate input features with opcode embeddings\n        \n        # Pass data through convolutional layers\n        for conv in self.convs:\n            x = conv(x, edge_index).relu()\n        \n        # Get 1D graph embedding using average pooling\n        x_graph = torch.mean(x, 0)\n        \n        # Combine graph data with config data\n        x = torch.cat([x_cfg, x_graph.repeat((len(x_cfg), 1))], axis=1)  # 🔄 Concatenate config data with repeated graph embeddings\n        \n        # Pass the combined data through the dense neural network\n        x = torch.flatten(self.dense(x))\n        \n        # Standardize the output\n        x = (x - torch.mean(x)) / (torch.std(x) + 1e-5)\n        return x\n\n# Create an instance of the 'SimpleModel' and move it to the specified device (CPU or GPU)\nmodel = SimpleModel(hidden_channels=[16, 32, 16, 48], graph_feats=64, hidden_dim=64).to(device)\n","metadata":{"papermill":{"duration":0.063336,"end_time":"2023-09-01T16:56:22.069894","exception":false,"start_time":"2023-09-01T16:56:22.006558","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 🚂 Train  Epoch","metadata":{"papermill":{"duration":0.008439,"end_time":"2023-09-01T16:56:22.088164","exception":false,"start_time":"2023-09-01T16:56:22.079725","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# 📊 Concatenate DataFrames\n# In this cell, we concatenate DataFrames 'train' and 'valid' from the 'tile_xla' dictionary along the row axis.\n# We then reset the index of the resulting DataFrame for consistent indexing.\n\n# Concatenate 'train' and 'valid' DataFrames along the row axis and reset the index\ndf = pd.concat((tile_xla[\"train\"], tile_xla[\"valid\"]), axis=0).reset_index(drop=True)\n\n# This operation combines the training and validation data for further processing, ensuring a unified DataFrame.\n","metadata":{"papermill":{"duration":0.021726,"end_time":"2023-09-01T16:56:22.11899","exception":false,"start_time":"2023-09-01T16:56:22.097264","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-09-13T18:47:55.200587Z","iopub.execute_input":"2023-09-13T18:47:55.200959Z","iopub.status.idle":"2023-09-13T18:47:55.210816Z","shell.execute_reply.started":"2023-09-13T18:47:55.200927Z","shell.execute_reply":"2023-09-13T18:47:55.209082Z"},"trusted":true},"execution_count":81,"outputs":[]},{"cell_type":"code","source":"# 🔄 Cross-Validation Training Loop (Enhanced)\n\n# Define the score_tile_mean function\ndef score_tile_mean(predictions, df):\n    score = 0\n    for i in range(len(df)):\n        predbest = np.mean(df.iloc[i]['config_runtime'][predictions[i]])\n        best = np.mean(np.sort(df.iloc[i]['config_runtime'])[:5])\n        score += 2 - predbest / best\n    score /= len(df)\n    return score\n\n# Define the score_tile_max function\ndef score_tile_max(predictions, df):\n    score = 0\n    for i in range(len(df)):\n        predbest = np.min(df.iloc[i]['config_runtime'][predictions[i]])\n        best = np.min(df.iloc[i]['config_runtime'])\n        score += 2 - predbest / best\n    score /= len(df)\n    return score\n\n# Create a K-Fold cross-validator with 5 splits\nkfold = sklearn.model_selection.KFold(n_splits=5, shuffle=True, random_state=0)\n\n# Lists to store mean and max scores for each fold\nscore_means = []\nscore_maxs = []\n\n# Define hyperparameters\nlearning_rate = 5e-4  # Adjust the learning rate to a different value\nweight_decay = 1e-6  # Adjust weight decay to a different value\nnum_epochs = 90  # You can keep the number of epochs as 90 or adjust as needed\n\n\n# Iterate through each fold\nfor fold, (tr_idx, va_idx) in enumerate(kfold.split(df)):\n    train_dataset = TileDataset(df.iloc[tr_idx])\n    val_dataset = TileDataset(df.iloc[va_idx])\n    criterion = torch.nn.MSELoss()\n    steps = len(train_dataset) * num_epochs  # Update the number of training steps\n    warmup_steps = int(steps * 0.1)\n    optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate, weight_decay=weight_decay)\n    scheduler = CosineLRScheduler(optimizer, t_initial=steps, warmup_t=warmup_steps, warmup_lr_init=1e-6, lr_min=2e-8)\n\n    best_score = 0\n    best_score_max = 0\n\n    # Training loop with increased epochs\n    for epoch in range(num_epochs):\n        model.train()\n        pbar = tqdm(range(len(train_dataset)), leave=False)\n        loss_sum = 0\n        n = 0\n        \n        for i in pbar:\n            cfg_ft, nd_ft, nd_op, ind, target = train_dataset[i]\n            cfg_ft, nd_ft, nd_op, ind, target = cfg_ft.to(device), nd_ft.to(device), nd_op.to(device), ind.to(device), target.to(device)\n            \n            out = model(cfg_ft, nd_ft, nd_op, ind)\n            loss = criterion(out, target)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1e-2)\n            scheduler.step(i + len(train_dataset) * epoch)\n            optimizer.step()\n            loss_sum += loss.item()\n            n += 1\n            pbar.set_description(f'running loss: {(loss_sum/n):.2f}, current loss: {(loss.item()):.2f}')\n        pbar.close()\n        model.eval()\n        tile_xla_predictions = []\n        pbar = tqdm(range(len(val_dataset)), leave=False)\n        \n        for i in pbar:\n            cfg_ft, nd_ft, nd_op, ind, target = val_dataset[i]\n            cfg_ft, nd_ft, nd_op, ind, target = cfg_ft.to(device), nd_ft.to(device), nd_op.to(device), ind.to(device), target.to(device)\n            \n            out = model(cfg_ft, nd_ft, nd_op, ind)\n            tile_xla_predictions.append(np.argsort(out.cpu().detach().numpy())[:5])\n        pbar.close()\n        \n        # Calculate and display scores for the current fold and epoch\n        score_mean = score_tile_mean(tile_xla_predictions, val_dataset.df)\n        score_max = score_tile_max(tile_xla_predictions, val_dataset.df)\n        print(f'fold {fold} epoch {epoch}, comp_score = {score_max:.3f}, mean_score = {score_mean:.3f},')\n        \n        # Update best scores and save the model if the mean score improves\n        if score_mean > best_score:\n            best_score = score_mean\n            best_score_max = score_max\n            torch.save(model.state_dict(), f'best_model_{fold}.pth')\n    \n    # Append the best scores for this fold to the respective lists\n    score_means.append(best_score)\n    score_maxs.append(best_score_max)\n\n# Calculate and display the mean scores across all folds\nprint(f'comp_score = {np.mean(score_maxs)}, mean_score = {np.mean(score_means)},')\n","metadata":{"papermill":{"duration":9138.052878,"end_time":"2023-09-01T19:28:40.180685","exception":false,"start_time":"2023-09-01T16:56:22.127807","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-09-13T18:48:12.310309Z","iopub.execute_input":"2023-09-13T18:48:12.310765Z"},"trusted":true},"execution_count":null,"outputs":[{"name":"stderr","text":"                                                                                              \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 0, comp_score = 0.608, mean_score = -0.014,\n","output_type":"stream"},{"name":"stderr","text":"                                                                                              \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 1, comp_score = 0.658, mean_score = 0.060,\n","output_type":"stream"},{"name":"stderr","text":"                                                                                              \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 2, comp_score = 0.592, mean_score = -0.039,\n","output_type":"stream"},{"name":"stderr","text":"                                                                                               \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 3, comp_score = 0.659, mean_score = 0.100,\n","output_type":"stream"},{"name":"stderr","text":"                                                                                           \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 4, comp_score = 0.785, mean_score = 0.565,\n","output_type":"stream"},{"name":"stderr","text":"                                                                                            \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 5, comp_score = 0.792, mean_score = 0.581,\n","output_type":"stream"},{"name":"stderr","text":"                                                                                            \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 6, comp_score = 0.885, mean_score = 0.777,\n","output_type":"stream"},{"name":"stderr","text":"                                                                                            \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 7, comp_score = 0.898, mean_score = 0.794,\n","output_type":"stream"},{"name":"stderr","text":"                                                                                              \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 10, comp_score = 0.891, mean_score = 0.738,\n","output_type":"stream"},{"name":"stderr","text":"running loss: 0.17, current loss: 0.18:   7%|▋         | 362/5108 [00:03<00:49, 95.69it/s]IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\nrunning loss: 0.23, current loss: 0.09:  78%|███████▊  | 3990/5108 [00:41<00:11, 97.16it/s] IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\n                                                                                            \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 21, comp_score = 0.926, mean_score = 0.839,\n","output_type":"stream"},{"name":"stderr","text":"running loss: 0.20, current loss: 0.30:  91%|█████████ | 4644/5108 [00:49<00:04, 98.66it/s] IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\nrunning loss: 0.20, current loss: 0.09:  68%|██████▊   | 3471/5108 [00:35<00:16, 100.40it/s]IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\n                                                                                            \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 25, comp_score = 0.924, mean_score = 0.836,\n","output_type":"stream"},{"name":"stderr","text":"                                                                                            \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 26, comp_score = 0.928, mean_score = 0.833,\n","output_type":"stream"},{"name":"stderr","text":"                                                                                            \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 28, comp_score = 0.933, mean_score = 0.841,\n","output_type":"stream"},{"name":"stderr","text":"running loss: 0.18, current loss: 0.09:  11%|█         | 572/5108 [00:05<00:42, 106.13it/s]IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\nrunning loss: 0.19, current loss: 0.19:  92%|█████████▏| 4703/5108 [00:45<00:03, 105.78it/s]IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\n                                                                                            \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 30, comp_score = 0.925, mean_score = 0.834,\n","output_type":"stream"},{"name":"stderr","text":"                                                       | 60/5108 [00:00<00:51, 98.03it/s]\r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 31, comp_score = 0.924, mean_score = 0.827,\n","output_type":"stream"},{"name":"stderr","text":"running loss: 0.19, current loss: 0.05:  53%|█████▎    | 2695/5108 [00:27<00:24, 99.51it/s] IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\n                                                                                            \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 32, comp_score = 0.927, mean_score = 0.835,\n","output_type":"stream"},{"name":"stderr","text":"running loss: 0.18, current loss: 0.10:  24%|██▍       | 1249/5108 [00:12<00:37, 102.29it/s]IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\n                                                                                            \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 33, comp_score = 0.922, mean_score = 0.830,\n","output_type":"stream"},{"name":"stderr","text":"running loss: 0.15, current loss: 0.15:   2%|▏         | 97/5108 [00:00<00:47, 104.60it/s]IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\nrunning loss: 0.19, current loss: 0.34:  84%|████████▍ | 4300/5108 [00:43<00:08, 94.98it/s] IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\nrunning loss: 0.19, current loss: 0.11:  56%|█████▌    | 2853/5108 [00:28<00:23, 97.68it/s] IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\n                                                                                            \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 35, comp_score = 0.916, mean_score = 0.821,\n","output_type":"stream"},{"name":"stderr","text":"running loss: 0.18, current loss: 0.06:  36%|███▌      | 1821/5108 [00:17<00:32, 102.15it/s]IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\n                                                                                            \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 36, comp_score = 0.927, mean_score = 0.838,\n","output_type":"stream"},{"name":"stderr","text":"running loss: 0.18, current loss: 0.06:  18%|█▊        | 922/5108 [00:08<00:39, 105.81it/s]IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\n                                                                                            \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 37, comp_score = 0.910, mean_score = 0.808,\n","output_type":"stream"},{"name":"stderr","text":"running loss: 0.18, current loss: 0.39:  31%|███       | 1576/5108 [00:16<00:35, 99.48it/s] IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\n                                                                                            \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 38, comp_score = 0.922, mean_score = 0.830,\n","output_type":"stream"},{"name":"stderr","text":"running loss: 0.16, current loss: 0.03:  14%|█▍        | 703/5108 [00:07<00:44, 98.99it/s] IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\nrunning loss: 0.18, current loss: 0.12:  98%|█████████▊| 4991/5108 [00:51<00:01, 96.38it/s] IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\nrunning loss: 0.18, current loss: 0.04:  79%|███████▊  | 4017/5108 [00:41<00:13, 81.20it/s] IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\nrunning loss: 0.19, current loss: 0.06:  61%|██████    | 3108/5108 [00:31<00:20, 98.70it/s] IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\n                                                                                            \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 41, comp_score = 0.918, mean_score = 0.826,\n","output_type":"stream"},{"name":"stderr","text":"running loss: 0.17, current loss: 0.11:  30%|██▉       | 1516/5108 [00:15<00:35, 101.37it/s]IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\n                                                                                            \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 42, comp_score = 0.928, mean_score = 0.836,\n","output_type":"stream"},{"name":"stderr","text":"running loss: 0.16, current loss: 0.04:   9%|▉         | 483/5108 [00:04<00:44, 102.89it/s]IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\nrunning loss: 0.18, current loss: 0.09:  94%|█████████▎| 4782/5108 [00:47<00:03, 102.41it/s]IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\nrunning loss: 0.18, current loss: 0.11:  73%|███████▎  | 3733/5108 [00:36<00:13, 104.98it/s]IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\n                                                                                            \r","output_type":"stream"},{"name":"stdout","text":"fold 0 epoch 44, comp_score = 0.928, mean_score = 0.838,\n","output_type":"stream"},{"name":"stderr","text":"running loss: 0.17, current loss: 0.03:  37%|███▋      | 1870/5108 [00:18<00:34, 93.91it/s] IOPub message rate exceeded.\nThe notebook server will temporarily stop sending output\nto the client in order to avoid crashing it.\nTo change this limit, set the config variable\n`--NotebookApp.iopub_msg_rate_limit`.\n\nCurrent values:\nNotebookApp.iopub_msg_rate_limit=1000.0 (msgs/sec)\nNotebookApp.rate_limit_window=3.0 (secs)\n\nrunning loss: 0.17, current loss: 0.10:  80%|████████  | 4094/5108 [01:03<00:15, 63.99it/s]]","output_type":"stream"}]},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":9138.052878,"end_time":"2023-09-01T19:28:40.180685","exception":false,"start_time":"2023-09-01T16:56:22.127807","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 📊 Evaluate on Validation Dataset","metadata":{"papermill":{"duration":20.903162,"end_time":"2023-09-01T19:29:21.465429","exception":false,"start_time":"2023-09-01T19:29:00.562267","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# 🚀 Predict and Submit (only tile:xla predictions)","metadata":{"papermill":{"duration":20.552961,"end_time":"2023-09-01T19:30:43.664106","exception":false,"start_time":"2023-09-01T19:30:23.111145","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# 📊 Predict on Test Dataset (tile:xla)\n# In this section, we use the trained model to make predictions on the test dataset ('tile:xla').\n\n# Create a TileDataset for the 'tile:xla' test dataset\ndataset = TileDataset(tile_xla[\"test\"])\n\n# List to store model predictions for each sample in the test dataset\ntile_xla_predictions = [[] for i in range(len(dataset))]\n\n# Iterate through each fold (previously trained models)\nfor fold in range(5):\n    # Load the trained model weights for the current fold\n    model.load_state_dict(torch.load(f'/kaggle/working/best_model_{fold}.pth'))\n    model.eval()  # 🕵️ Set the model to evaluation mode\n    pbar = tqdm(range(len(dataset)))  # Progress bar for test data prediction\n    \n    for i in pbar:\n        cfg_ft, nd_ft, nd_op, ind, target = dataset[i]\n        cfg_ft, nd_ft, nd_op, ind, target = cfg_ft.to(device), nd_ft.to(device), nd_op.to(device), ind.to(device), target.to(device)\n\n        out = model(cfg_ft, nd_ft, nd_op, ind)\n        tile_xla_predictions[i].append(out.cpu().detach().numpy())\n\n# Aggregate predictions by taking the mean and selecting the top 5\ntile_xla_predictions = [np.argsort(np.mean(pred, axis=0))[:5] for pred in tile_xla_predictions]\n\n# The 'tile_xla_predictions' now contains the top 5 predicted results for each sample in the 'tile:xla' test dataset.\ntile_xla_predictions\n","metadata":{"papermill":{"duration":42.864288,"end_time":"2023-09-01T19:31:46.765649","exception":false,"start_time":"2023-09-01T19:31:03.901361","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 📊 Generate and Save Submission File\n# In this section, we generate a submission file based on the model predictions and save it.\n\n# Read the sample submission file\nsub = pd.read_csv('/kaggle/input/predict-ai-model-runtime/sample_submission.csv')\n\n# Iterate through the test file names and update the submission file with top predictions\nfor i, filename in enumerate(tile_xla[\"test\"]['file'].values):\n    id = 'tile:xla:' + filename[:-4]  # Construct the ID for the submission\n    sub.loc[sub.ID == id, 'TopConfigs'] = ';'.join(tile_xla_predictions[i].astype(str))\n\n# Save the updated submission file as 'submission.csv' without the index\nsub.to_csv('submission.csv', index=False)\n\n# Display the updated submission file\nsub\n","metadata":{"papermill":{"duration":20.880172,"end_time":"2023-09-01T19:32:28.307392","exception":false,"start_time":"2023-09-01T19:32:07.42722","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]}]}