{"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"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":58266,"databundleVersionId":6641124,"sourceType":"competition"},{"sourceId":6991586,"sourceType":"datasetVersion","datasetId":3931672},{"sourceId":144043388,"sourceType":"kernelVersion"},{"sourceId":144045966,"sourceType":"kernelVersion"},{"sourceId":144045983,"sourceType":"kernelVersion"}],"dockerImageVersionId":30558,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# BERT Tile model with [5 fold] cross validation and GNN Layout model","metadata":{"_uuid":"38065b54-381a-4921-a43f-ce1b68fa83c6","_cell_guid":"1fae0bfd-fb6b-4f81-9eec-cd66beeca4d5","trusted":true}},{"cell_type":"code","source":"!python --version\n!pip install --upgrade pip","metadata":{"_uuid":"14046626-a947-4ea3-b569-37dd16064ef3","_cell_guid":"a6109b74-a91a-4171-8a7a-1420aac8be01","scrolled":true,"execution":{"iopub.status.busy":"2023-11-20T03:51:44.584101Z","iopub.execute_input":"2023-11-20T03:51:44.584647Z","iopub.status.idle":"2023-11-20T03:52:00.032466Z","shell.execute_reply.started":"2023-11-20T03:51:44.584601Z","shell.execute_reply":"2023-11-20T03:52:00.031287Z"},"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -U joblib\n!pip install -U datasets\n!pip install -U optuna\n!pip install -U ipywidgets","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-11-20T03:52:00.035215Z","iopub.execute_input":"2023-11-20T03:52:00.035603Z","iopub.status.idle":"2023-11-20T03:52:57.894941Z","shell.execute_reply.started":"2023-11-20T03:52:00.035572Z","shell.execute_reply":"2023-11-20T03:52:57.893437Z"},"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import datetime, os, time, sys\nfrom pathlib import Path\nfrom typing import Dict, Optional, List, Union, Tuple\nfrom dataclasses import dataclass\nimport math\nimport numpy as np\nimport pandas as pd\nfrom datasets import Dataset\nimport time, gc\nimport joblib\nimport optuna\nfrom tqdm import tqdm","metadata":{"_uuid":"135d0801-1cdc-48aa-a835-a56bcf81e4ec","_cell_guid":"2d8066b1-ea85-4958-b4a8-de9c496ab4d5","execution":{"iopub.status.busy":"2023-11-20T03:52:57.896987Z","iopub.execute_input":"2023-11-20T03:52:57.897366Z","iopub.status.idle":"2023-11-20T03:52:57.906243Z","shell.execute_reply.started":"2023-11-20T03:52:57.897334Z","shell.execute_reply":"2023-11-20T03:52:57.904724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Layout inference with Torch GNN\n","metadata":{}},{"cell_type":"code","source":"!pip install -U dill","metadata":{"execution":{"iopub.status.busy":"2023-11-20T03:52:57.908961Z","iopub.execute_input":"2023-11-20T03:52:57.909469Z","iopub.status.idle":"2023-11-20T03:53:12.201367Z","shell.execute_reply.started":"2023-11-20T03:52:57.909426Z","shell.execute_reply":"2023-11-20T03:53:12.200238Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -U tensorflow_gnn --pre\n!pip install -U tensorflow_ranking","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-11-20T03:53:12.204966Z","iopub.execute_input":"2023-11-20T03:53:12.205942Z","iopub.status.idle":"2023-11-20T03:53:42.209655Z","shell.execute_reply.started":"2023-11-20T03:53:12.205905Z","shell.execute_reply":"2023-11-20T03:53:42.208173Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tensorflow as tf\nimport tensorflow_gnn as tfgnn\nimport tensorflow_ranking as tfr\n\nimport tpugraphsv1_layout_data_py as layout_data\nimport tpugraphsv1_tile_data_py as tile_data\nimport tpugraphsv1_implicit_py as implicit\n\nprint(tf.__version__)","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-11-20T03:53:42.211682Z","iopub.execute_input":"2023-11-20T03:53:42.212061Z","iopub.status.idle":"2023-11-20T03:53:42.218963Z","shell.execute_reply.started":"2023-11-20T03:53:42.212031Z","shell.execute_reply":"2023-11-20T03:53:42.217634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## ResModel for layout configs\n","metadata":{}},{"cell_type":"code","source":"def _mlp(dims, hidden_activation, l2reg=1e-4, use_bias=True):\n    \"\"\"Helper function for multi-layer perceptron (MLP).\"\"\"\n    layers = []\n    for i, dim in enumerate(dims):\n        if i > 0:\n            layers.append(tf.keras.layers.Activation(hidden_activation))\n        layers.append(tf.keras.layers.Dense(dim, kernel_regularizer=tf.keras.regularizers.l2(l2reg),\n                                            use_bias=use_bias))\n    return tf.keras.Sequential(layers)\n\n\nclass _OpEmbedding(tf.keras.Model):\n    \"\"\"Embeds GraphTensor.node_sets['op']['op'] nodes into feature 'op_e'.\"\"\"\n\n    def __init__(self, num_ops: int, embed_d: int, l2reg: float = 1e-4):\n        super().__init__()\n        self.embedding_layer = tf.keras.layers.Embedding(num_ops,\n                                                         embed_d,\n                                                         activity_regularizer=tf.keras.regularizers.l2(l2reg))\n\n    def call(self, graph: tfgnn.GraphTensor, training: bool = False) -> tfgnn.GraphTensor:\n        op_features = dict(graph.node_sets['op'].features)\n        op_features['op_e'] = self.embedding_layer(tf.cast(graph.node_sets['op']['op'], tf.int32))\n        return graph.replace_features(node_sets={'op': op_features})\n\n\ndef pair_layout_graph_with_label(graph: tfgnn.GraphTensor):\n    \"\"\"Extracts label from graph (`tfgnn.GraphTensor`) and returns a pair of `(graph, label)`\"\"\"\n    # Return runtimes divded over large number: only ranking is required. The\n    # runtimes are in the 100K range\n    label = tf.cast(graph.node_sets['g']['runtimes'], tf.float32) / 1e7\n    return graph, label","metadata":{"execution":{"iopub.status.busy":"2023-11-20T03:53:42.221247Z","iopub.execute_input":"2023-11-20T03:53:42.221697Z","iopub.status.idle":"2023-11-20T03:53:42.237260Z","shell.execute_reply.started":"2023-11-20T03:53:42.221658Z","shell.execute_reply":"2023-11-20T03:53:42.235933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ResModel(tf.keras.Model):\n    \"\"\"GNN with residual connections.\"\"\"\n    def __init__(self, num_configs: int, num_ops: int, op_embed_dim: int = 32,\n                 num_gnns: int = 2, mlp_layers: int = 2,\n                 hidden_activation: str = 'leaky_relu',\n                 hidden_dim: int = 32, reduction: str = 'sum'):\n        super().__init__()\n        self._num_configs = num_configs\n        self._num_ops = num_ops\n        self._op_embedding = _OpEmbedding(num_ops, op_embed_dim)\n        self._prenet = _mlp([hidden_dim] * mlp_layers, hidden_activation)\n        self._gc_layers = []\n        for _ in range(num_gnns):\n            self._gc_layers.append(_mlp([hidden_dim] * mlp_layers, hidden_activation))\n        self._postnet = _mlp([hidden_dim, 1], hidden_activation, use_bias=False)\n\n    def call(self, graph: tfgnn.GraphTensor, training: bool = False):\n        del training\n        return self.forward(graph, self._num_configs)\n\n    def _node_level_forward(self, node_features: tf.Tensor,\n                            config_features: tf.Tensor,\n                            graph: tfgnn.GraphTensor, num_configs: int,\n                            edgeset_prefix='') -> tf.Tensor:\n        adj_op_op = implicit.AdjacencyMultiplier(graph, edgeset_prefix+'feed')  # op->op\n        adj_config = implicit.AdjacencyMultiplier(graph, edgeset_prefix+'config')  # nconfig->op\n\n        adj_op_op_hat = (adj_op_op + adj_op_op.transpose()).add_eye()\n        adj_op_op_hat = adj_op_op_hat.normalize_symmetric()\n\n        x = node_features\n\n        x = tf.stack([x] * num_configs, axis=1)\n        config_features = 100 * (adj_config @ config_features)\n        x = tf.concat([config_features, x], axis=-1)\n        x = self._prenet(x)\n        x = tf.nn.leaky_relu(x)\n\n        for layer in self._gc_layers:\n            y = x\n            y = tf.concat([config_features, y], axis=-1)\n            y = tf.nn.leaky_relu(layer(adj_op_op_hat @ y))\n            x += y\n        return x\n\n    def forward(self, graph: tfgnn.GraphTensor, num_configs: int, backprop=True) -> tf.Tensor:\n        graph = self._op_embedding(graph)\n\n        config_features = graph.node_sets['nconfig']['feats']\n        node_features = tf.concat([ graph.node_sets['op']['feats'],\n                                    graph.node_sets['op']['op_e']], axis=-1)\n\n        x_full = self._node_level_forward(node_features=tf.stop_gradient(node_features),\n                                          config_features=tf.stop_gradient(config_features),\n                                          graph=graph, num_configs=num_configs)\n\n        if backprop:\n            x_backprop = self._node_level_forward(\n                                node_features=node_features,\n                                config_features=config_features,\n                                graph=graph, num_configs=num_configs,\n                                edgeset_prefix='sampled_')\n\n            is_selected = graph.node_sets['op']['selected']\n            # Need to expand twice as `is_selected` is a vector (num_nodes) but\n            # x_{backprop, full} are 3D tensors (num_nodes, num_configs, num_feats).\n            is_selected = tf.expand_dims(is_selected, axis=-1)\n            is_selected = tf.expand_dims(is_selected, axis=-1)\n            x = tf.where(is_selected, x_backprop, x_full)\n        else:\n            x = x_full\n        # Multiplication of adjacency matrix of two graphs\n        adj_config = implicit.AdjacencyMultiplier(graph, 'config')\n\n        # Features for configurable nodes.\n        config_feats = (adj_config.transpose() @ x)\n\n        # Global pooling\n        adj_pool_op_sum = implicit.AdjacencyMultiplier(graph, 'g_op').transpose()\n        adj_pool_op_mean = adj_pool_op_sum.normalize_right()\n        adj_pool_config_sum = implicit.AdjacencyMultiplier(graph, 'g_config').transpose()\n        x = self._postnet(tf.concat([\n            # (A D^-1) @ Features\n            adj_pool_op_mean @ x,\n            # l2_normalize( A @ Features )\n            tf.nn.l2_normalize(adj_pool_op_sum @ x, axis=-1),\n            # l2_normalize( A @ Features )\n            tf.nn.l2_normalize(adj_pool_config_sum @ config_feats, axis=-1),\n        ], axis=-1))\n\n        x = tf.squeeze(x, -1)\n\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-11-20T03:53:42.239287Z","iopub.execute_input":"2023-11-20T03:53:42.239718Z","iopub.status.idle":"2023-11-20T03:53:42.263894Z","shell.execute_reply.started":"2023-11-20T03:53:42.239686Z","shell.execute_reply":"2023-11-20T03:53:42.262712Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training layout model using optuma libary\n\nWe use optuma library to tune the hyperparameters of the model","metadata":{}},{"cell_type":"code","source":"AUTO = tf.data.experimental.AUTOTUNE\n\n# Detect TPU, return appropriate distribution strategy\ntry:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver() \n    print('Running on TPU ', tpu.master())\nexcept ValueError:\n    tpu = None\n\nif tpu:\n    tf.config.experimental_connect_to_cluster(tpu)\n    tf.tpu.experimental.initialize_tpu_system(tpu)\n    strategy = tf.distribute.experimental.TPUStrategy(tpu)\nelse:\n    strategy = tf.distribute.get_strategy() ","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-11-20T03:53:42.265306Z","iopub.execute_input":"2023-11-20T03:53:42.265721Z","iopub.status.idle":"2023-11-20T03:53:42.279885Z","shell.execute_reply.started":"2023-11-20T03:53:42.265668Z","shell.execute_reply":"2023-11-20T03:53:42.278539Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Batch size information.\nBATCH_SIZE = 16 * strategy.num_replicas_in_sync # Number of graphs per batch.\n\nCONFIGS_PER_GRAPH = 5  # Number of configurations (features and target values) per graph.\nMAX_NUM_CONFIGS = 500 # Maximal Number of configurations used for filter. Default value = 500 \nMAX_KEEP_NODES = 1000  # Useful for dropout.\nBUFFER_SIZE = 10000\n\nprint(BATCH_SIZE)","metadata":{"execution":{"iopub.status.busy":"2023-11-20T03:53:42.281180Z","iopub.execute_input":"2023-11-20T03:53:42.281553Z","iopub.status.idle":"2023-11-20T03:53:42.291647Z","shell.execute_reply.started":"2023-11-20T03:53:42.281521Z","shell.execute_reply":"2023-11-20T03:53:42.290585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def clear_memory():\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-11-20T03:53:42.292967Z","iopub.execute_input":"2023-11-20T03:53:42.293313Z","iopub.status.idle":"2023-11-20T03:53:42.301317Z","shell.execute_reply.started":"2023-11-20T03:53:42.293284Z","shell.execute_reply":"2023-11-20T03:53:42.300441Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Split into train and valid layout dataset","metadata":{}},{"cell_type":"code","source":"layout_npz_dataset = None  # declare as global","metadata":{"execution":{"iopub.status.busy":"2023-11-20T03:53:42.302780Z","iopub.execute_input":"2023-11-20T03:53:42.303082Z","iopub.status.idle":"2023-11-20T03:53:42.314552Z","shell.execute_reply.started":"2023-11-20T03:53:42.303055Z","shell.execute_reply":"2023-11-20T03:53:42.313163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# `MAX_KEEP_NODES` is (or, is not) useful for Segment Dropout, if model uses\n# edges \"sampled_config\" and \"sampled_feed\" (or, \"config\" and \"feed\")\ndef split_layout_dataset(source, search):\n    global layout_npz_dataset\n    layout_data_root_dir = os.path.join(os.path.expanduser(LAYOUT_DATA_ROOT), source, search)\n    # Load layout dataset\n    layout_npz_dataset = layout_data.get_npz_dataset(\n        layout_data_root_dir,\n        min_train_configs=CONFIGS_PER_GRAPH,\n        max_train_configs=MAX_NUM_CONFIGS,  # Default: 500 If any graph has more than this configurations, it will be filtered [speeds up loading + training]\n        cache_dir=f'cache_{source}_{search}'\n    )\n    ## Layout training dataset\n    layout_train_ds = (layout_npz_dataset.train.get_graph_tensors_dataset(CONFIGS_PER_GRAPH, max_nodes=MAX_KEEP_NODES)\n                            .shuffle(100, reshuffle_each_iteration=True)\n                            .batch(BATCH_SIZE, drop_remainder=False)\n                            .map(tfgnn.GraphTensor.merge_batch_to_components)\n                            .map(pair_layout_graph_with_label))\n    # Layout valid dataset\n    layout_valid_ds = (layout_npz_dataset.validation.get_graph_tensors_dataset(CONFIGS_PER_GRAPH)\n                            .batch(BATCH_SIZE, drop_remainder=False) \n                            .map(tfgnn.GraphTensor.merge_batch_to_components)\n                            .map(pair_layout_graph_with_label))\n                       \n#     print(next(iter(layout_train_ds)))\n    return layout_train_ds, layout_valid_ds","metadata":{"execution":{"iopub.status.busy":"2023-11-20T03:53:42.316153Z","iopub.execute_input":"2023-11-20T03:53:42.316875Z","iopub.status.idle":"2023-11-20T03:53:42.326067Z","shell.execute_reply.started":"2023-11-20T03:53:42.316840Z","shell.execute_reply":"2023-11-20T03:53:42.324849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Build the layout model ","metadata":{}},{"cell_type":"code","source":"def build_model(num_ops):    \n    # Create a ResModel\n    model = ResModel(CONFIGS_PER_GRAPH, num_ops)\n\n    loss = tfr.keras.losses.ListMLELoss()  # (temperature=10)\n    # opt = tf.keras.optimizers.Adam(learning_rate=1e-3, clipnorm=0.5)\n    opt = tf.keras.optimizers.AdamW(learning_rate=1e-3, clipnorm=0.5)\n    # opt = tf.keras.optimizers.SGD(lr=1e-3)\n\n    model.compile(loss=loss, optimizer=opt,\n                  metrics=[tfr.keras.metrics.OPAMetric(name='opa_metric')],\n                  steps_per_execution=32)\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-11-20T03:53:42.332585Z","iopub.execute_input":"2023-11-20T03:53:42.333333Z","iopub.status.idle":"2023-11-20T03:53:42.340201Z","shell.execute_reply.started":"2023-11-20T03:53:42.333298Z","shell.execute_reply":"2023-11-20T03:53:42.339095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training the layout model","metadata":{}},{"cell_type":"code","source":"def train_model_by_epoch(epoch, model, layout_train_ds, layout_valid_ds):\n    global best_params, best_val_opa, best_val_at_epoch, early_stop\n    print(f\"Starting training the model with EPOCH {epoch}\")\n    start = time.time()\n\n    # Train the model \n    history = model.fit(layout_train_ds, epochs=epoch+1, verbose=\"auto\", \n                        batch_size = BATCH_SIZE,\n                        workers=4, validation_data=layout_valid_ds) # Frequency of validation\n\n    # get the training result\n    print(f\"epoch = {epoch} history = {history.history}\")\n#             train_loss = history.history['loss'][-1]\n#             train_opa = history.history['opa_metric'][-1]\n#             val_loss = history.history['val_loss'][-1]\n    val_opa = history.history['val_opa_metric'][-1]\n    if val_opa > best_val_opa:\n        best_val_opa = val_opa\n        best_val_at_epoch = epoch\n        best_params = {v.ref: v + 0 for v in model.trainable_variables}\n        print(f' * [{epoch}] Validation (NEW BEST): {val_opa}')\n        \n    print(f\"Finish the training for EPOCH {epoch} in {time.time() - start}\")","metadata":{"execution":{"iopub.status.busy":"2023-11-20T03:53:42.341647Z","iopub.execute_input":"2023-11-20T03:53:42.342626Z","iopub.status.idle":"2023-11-20T03:53:42.358106Z","shell.execute_reply.started":"2023-11-20T03:53:42.342592Z","shell.execute_reply":"2023-11-20T03:53:42.357022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"LAYOUT_DATA_ROOT = '/kaggle/input/predict-ai-model-runtime/npz_all/npz/layout'\nTOTAL_EPOCHS = 1 # Total number of epochs\ndef training_model(source, search, is_pretrain=True):\n    start = time.time()\n    # Split layout dataset into 'train' and 'valid' datasets\n    layout_train_ds, layout_valid_ds = split_layout_dataset(source, search)\n    num_ops = layout_npz_dataset.num_ops\n    # Create a ResModel\n    model = build_model(num_ops)\n    try:            \n        # This initializes the variables used by the optimizers,\n        # as well as any stateful metric variables        \n        model.fit(layout_train_ds, epochs=0, verbose=\"auto\")\n        model.load_weights(f'/kaggle/input/bert-like-title-model-output/layout_{source}_{search}/best_model_{source}_{search}')\n    except Exception as error:\n        print(\"An exception occurred during model loading \", error)\n    \n#     if is_pretrain == True: \n#         return model\n    \n    # ### Train for a few epochs.\n    global best_params, best_val_opa, best_val_at_epoch, early_stop \n    early_stop = 5  # If validation OPA did not increase in this many epochs, terminate training.\n    best_params = None  # Stores parameters corresponding to best validation OPA, to restore to them after training.\n    best_val_opa = -1  # Tracks best validation OPA\n    best_val_at_epoch = -1  # At which epoch.\n    epochs = TOTAL_EPOCHS # Total number of training epochs.\n    \n    for epoch in range(epochs):\n        try:\n            train_model_by_epoch(epoch, model, layout_train_ds, layout_valid_ds)\n            if early_stop > 0 and epoch - best_val_at_epoch >= early_stop:\n                print(f'[{epoch}] Best accuracy was attained at epoch {best_val_at_epoch}. Stopping.')\n                break\n        except Exception as error:\n            print(\"An exception occurred during trainig the model \", error)\n    # Restore best parameters.\n    print('Restoring parameters corresponding to the best validation OPA.')\n    assert best_params is not None\n    for v in model.trainable_variables:\n        v.assign(best_params[v.ref])\n        \n    model.save_weights(f'/kaggle/working/layout_{source}_{search}/best_model_{source}_{search}')\n    \n    del layout_train_ds, layout_valid_ds\n    print(f\"Training time of {source}-{search}: {time.time() - start}\")\n    return model","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-11-20T03:53:42.360317Z","iopub.execute_input":"2023-11-20T03:53:42.361010Z","iopub.status.idle":"2023-11-20T03:53:42.375358Z","shell.execute_reply.started":"2023-11-20T03:53:42.360967Z","shell.execute_reply":"2023-11-20T03:53:42.374451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Infer the layouts of test dataset","metadata":{}},{"cell_type":"code","source":"# Take a cut of the configs.\ndef infer_worker(i, graph, model, num_configs):\n    end_i = min(i + INFERENCE_CONFIGS_BATCH_SIZE, num_configs)\n    # Take a cut of the configs.\n    node_set_g = graph.node_sets['g']\n    subconfigs_graph = tfgnn.GraphTensor.from_pieces(\n        edge_sets=graph.edge_sets, ## Edges\n        ## Node set\n        node_sets={'op': graph.node_sets['op'],\n                   'nconfig': tfgnn.NodeSet.from_fields(\n                               sizes=graph.node_sets['nconfig'].sizes,\n                               features={'feats': graph.node_sets['nconfig']['feats'][:, i:end_i]}\n                                                       ),\n                   'g': tfgnn.NodeSet.from_fields(\n                               sizes=tf.constant([1]),\n                               features={'graph_id': node_set_g['graph_id'],\n                                         'runtimes': node_set_g['runtimes'][:, i:end_i],\n                                         'kept_node_ratio': node_set_g['kept_node_ratio']}\n                                                  ) # End of 'g'\n\n                    } # End of 'node_sets'\n        )\n    h = model.forward(subconfigs_graph, num_configs=(end_i - i), backprop=False) # Don't update model weights\n#     return h[0]\n    global all_scores  # needed to modify the global value\n    all_scores.append(h[0]) \n    # print(f\"Total number of scores = {len(all_scores)}\")\n    ","metadata":{"execution":{"iopub.status.busy":"2023-11-20T03:53:42.377089Z","iopub.execute_input":"2023-11-20T03:53:42.378425Z","iopub.status.idle":"2023-11-20T03:53:42.390758Z","shell.execute_reply.started":"2023-11-20T03:53:42.378364Z","shell.execute_reply":"2023-11-20T03:53:42.389538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import concurrent\nfrom concurrent.futures import ThreadPoolExecutor\nfrom functools import partial\n\nINFERENCE_CONFIGS_BATCH_SIZE = 50\n# Source can be \"xla\" or \"nlp\". Search can be \"random\" or \"default\"\ndef infer_layout(source, search, is_pretrain=True):\n    # Training the model\n    model = training_model(source, search, is_pretrain)\n    # Infer the results using the model\n    start = time.time()\n    # Create the submission file\n    output_csv_filename = f'inference_layout_{source}_{search}.csv'\n    print('\\n\\n   Running inference on test set ...\\n\\n')\n    # Store the results of test dataset\n    test_rankings = []\n    assert layout_npz_dataset.test.graph_id is not None\n    for graph in tqdm(layout_npz_dataset.test.iter_graph_tensors(),\n                      total=layout_npz_dataset.test.graph_id.shape[-1],\n                      desc='Inference'):\n        num_configs = graph.node_sets['g']['runtimes'].shape[-1]\n        print(f\"num_configs = {num_configs}\")\n        global all_scores # declare as a global value\n        all_scores = []\n#         func = partial(infer_worker, graph=graph, model=model, num_configs=num_configs)\n#         with concurrent.futures.ThreadPoolExecutor(max_workers=3) as executor:\n#             # execute tasks concurrently and process results in order\n#             for result in executor.map(func, list(range(0, num_configs, INFERENCE_CONFIGS_BATCH_SIZE))):\n#                 all_scores.append(result) \n        for i in tqdm(range(0, num_configs, INFERENCE_CONFIGS_BATCH_SIZE)):\n            infer_worker(i, graph, model, num_configs)\n        all_scores = tf.concat(all_scores, axis=0)\n        graph_id = graph.node_sets['g']['graph_id'][0].numpy().decode()\n        sorted_indices = tf.strings.join(tf.strings.as_string(tf.argsort(all_scores)), ';').numpy().decode()\n        test_rankings.append((f\"layout:{source}:{search}:{graph_id}\", sorted_indices))\n    # Write the test_ranking \n    df = pd.DataFrame(test_rankings, columns=['ID', 'TopConfigs'])\n    df.to_csv(output_csv_filename)\n\n    del model\n    clear_memory()\n    print(f\"Total inference time of {source}-{search}: {time.time() - start}\")\n    return test_rankings","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-11-20T03:53:42.392403Z","iopub.execute_input":"2023-11-20T03:53:42.392763Z","iopub.status.idle":"2023-11-20T03:53:42.408217Z","shell.execute_reply.started":"2023-11-20T03:53:42.392735Z","shell.execute_reply":"2023-11-20T03:53:42.407060Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train and Run Layout Model","metadata":{}},{"cell_type":"code","source":"IS_DEBUG = False\nif IS_DEBUG:\n    #Load layout results\n    layout_test_df = pd.read_csv('/kaggle/input/bert-like-title-model-output/inference_layout_all.csv') \nelse:\n    test_rankings = []\n    test_rankings.extend(infer_layout('nlp', 'default'))\n    test_rankings.extend(infer_layout('nlp', 'random'))\n    test_rankings.extend(infer_layout('xla', 'default'))\n    test_rankings.extend(infer_layout('xla', 'random'))\n\n    layout_test_df = pd.DataFrame(test_rankings, columns=['ID', 'TopConfigs'])    \n    # Save to inference_layout_all.csv\n    layout_test_df.to_csv('/kaggle/working/inference_layout_all.csv')\n    layout_test_df.head(3)","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-11-20T03:53:42.409722Z","iopub.execute_input":"2023-11-20T03:53:42.410090Z","iopub.status.idle":"2023-11-20T03:53:42.439139Z","shell.execute_reply.started":"2023-11-20T03:53:42.410058Z","shell.execute_reply":"2023-11-20T03:53:42.438223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Tile model","metadata":{"_uuid":"80ba1a09-1208-49a2-a698-a069635a3ea6","_cell_guid":"4a0d5139-3717-4ba5-aaad-40273b344263","trusted":true}},{"cell_type":"code","source":"!pip install -U transformers\n!pip install -U pytorch-lightning","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-11-20T03:53:42.440820Z","iopub.execute_input":"2023-11-20T03:53:42.441530Z","iopub.status.idle":"2023-11-20T03:54:11.375239Z","shell.execute_reply.started":"2023-11-20T03:53:42.441497Z","shell.execute_reply":"2023-11-20T03:54:11.373788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom torch.nn.utils.rnn import pad_sequence\nfrom torch.utils.data import DataLoader\nfrom torch.optim import Adam, SGD, AdamW\nimport torchmetrics as tm\n\nfrom transformers.modeling_outputs import BaseModelOutputWithPastAndCrossAttentions\nfrom transformers.pytorch_utils import apply_chunking_to_forward\nfrom transformers.activations import ACT2FN\nimport pytorch_lightning as pl\nprint(pl.__version__)","metadata":{"execution":{"iopub.status.busy":"2023-11-20T03:54:11.377332Z","iopub.execute_input":"2023-11-20T03:54:11.378126Z","iopub.status.idle":"2023-11-20T03:54:11.387615Z","shell.execute_reply.started":"2023-11-20T03:54:11.378081Z","shell.execute_reply":"2023-11-20T03:54:11.386352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NODE_OP_CODES = 120\nNODE_FEATS = 140\nCONFIG_FEATS = 24\nNODE_CONFIG_FEATS = 18","metadata":{"execution":{"iopub.status.busy":"2023-11-20T03:54:11.389118Z","iopub.execute_input":"2023-11-20T03:54:11.389469Z","iopub.status.idle":"2023-11-20T03:54:11.400115Z","shell.execute_reply.started":"2023-11-20T03:54:11.389441Z","shell.execute_reply":"2023-11-20T03:54:11.399076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset\n\nCreate an Adjacency matrix for masking the attention\nCreates a virtual first node equivalent to the [CLS] token which contains the global config for tile cases, while layout node configuration goes to the corresponding node position","metadata":{}},{"cell_type":"code","source":"DATA_DIR = \"../input/predict-ai-model-runtime/npz_all/npz\"\n\ndef generate_tile_df() -> pd.DataFrame:\n    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        ID=lambda df: df.collection + ':' + df.model_name ,\n        paths = lambda df: df.paths.apply(lambda x: str(x))\n    )\n    return tile_df","metadata":{"execution":{"iopub.status.busy":"2023-11-20T03:54:11.402165Z","iopub.execute_input":"2023-11-20T03:54:11.402665Z","iopub.status.idle":"2023-11-20T03:54:11.413055Z","shell.execute_reply.started":"2023-11-20T03:54:11.402622Z","shell.execute_reply":"2023-11-20T03:54:11.411842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tile_df = generate_tile_df()\ntile_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-11-20T03:54:11.415092Z","iopub.execute_input":"2023-11-20T03:54:11.415591Z","iopub.status.idle":"2023-11-20T03:54:29.218575Z","shell.execute_reply.started":"2023-11-20T03:54:11.415553Z","shell.execute_reply":"2023-11-20T03:54:29.217263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def edges_adjacency(edges: torch.Tensor, add_diagonal=True) -> torch.Tensor:\n    \"\"\"\n    Generate an adjacency matrix from the edges\n    Args:\n        edges: Tensor of shape (num_edges, 2) with the edges\n        add_diagonal: Boolean indicating if the diagonal should be added to the adjacency matrix\n    Returns:\n        adjacency_matrix: Tensor of shape (num_nodes, num_nodes) with the adjacency matrix\n    \"\"\"\n    adjacency_matrix = torch.zeros((edges.max() + 1, edges.max() + 1))\n    adjacency_matrix[edges[:, 0], edges[:, 1]] = 1\n    if add_diagonal:\n        diag_idx = torch.arange(adjacency_matrix.shape[0])\n        adjacency_matrix[diag_idx, diag_idx] = 1\n    return adjacency_matrix\n\ndef node_cls_token(elem_dict, shift_node_config_ids:bool=True):\n    \"\"\"\n    Add a cls token to the node opcode, features, edges adjacency matrix, shift node_config_ids by 1 to account for the cls token\n    Args:\n        elem_dict: Dictionary with the elements of the tile\n    Returns:\n        elem_dict: Dictionary with the elements of the tile with the cls token\n    \"\"\"\n    elem_dict['node_opcode'] = torch.cat([torch.tensor([0]), elem_dict['node_opcode']]) # Introduce [CLS] node\n    elem_dict['node_feat'] = torch.cat([torch.zeros((1, elem_dict['node_feat'].shape[1])), elem_dict['node_feat']])\n    elem_dict['edges_adjecency'] = F.pad(elem_dict['edges_adjecency'], (1,0,1,0), value=1)\n    if 'node_config_ids' in elem_dict and shift_node_config_ids:\n        elem_dict['node_config_ids'] = elem_dict['node_config_ids'] + 1 # Shift Node Config IDs to take in to account [CLS] node\n    return elem_dict\n\n\nclass TileDataset(torch.utils.data.Dataset):\n    \n    def __init__(self, df:pd.DataFrame ,add_cls_token:bool=True,\n                 num_configs:int=10,  max_configs:Optional[int]=None):\n        self.df = df\n        self.add_cls_token = add_cls_token\n        self.num_configs = num_configs\n        self.max_configs = max_configs  \n        \n    def __len__(self) -> int:\n        return len(self.df)\n    \n    def select_configs(self, total_configs:int):\n        if self.max_configs is not None:\n            total_configs = min(total_configs, self.max_configs)\n        if self.num_configs == -1:\n            return np.arange(total_configs)\n        if total_configs < self.num_configs:\n            return np.random.choice(total_configs, self.num_configs, replace=True)\n        return  np.random.choice(total_configs, self.num_configs, replace=False)\n    \n    def tile_loader(self, 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['edge_index'])\n        return tile_dict\n    \n    \n    def __getitem__(self, idx:int, selected_configs:List[int]=None):\n        tile_dict = self.tile_loader(self.df.paths[idx])\n        if selected_configs is None:\n            selected_configs = self.select_configs(tile_dict['config_feat'].shape[0])\n        tile_dict['node_config_feat'] = tile_dict.pop('config_feat')[selected_configs]\n        tile_dict['node_config_feat'] = F.pad(tile_dict['node_config_feat'].unsqueeze(1), (0,NODE_CONFIG_FEATS))\n        tile_dict['config_runtime'] = tile_dict['config_runtime'][selected_configs].float()\n        tile_dict['config_runtime'] /= tile_dict['config_runtime_normalizers'][selected_configs].float()\n        tile_dict['node_config_ids'] = torch.zeros((1,))\n        tile_dict['selected_idxs'] = selected_configs\n        if self.add_cls_token:\n            tile_dict = node_cls_token(tile_dict, False)\n        return tile_dict\n","metadata":{"_uuid":"fe562304-0992-45d7-b8a7-46901fa4b24a","_cell_guid":"5c83e1c3-53b0-4089-99b6-0f0e140f5328","execution":{"iopub.status.busy":"2023-11-20T03:54:29.220536Z","iopub.execute_input":"2023-11-20T03:54:29.220949Z","iopub.status.idle":"2023-11-20T03:54:29.243571Z","shell.execute_reply.started":"2023-11-20T03:54:29.220909Z","shell.execute_reply":"2023-11-20T03:54:29.242270Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tile_dataset = TileDataset(tile_df)\ntile_elem = tile_dataset[0]\nfor k,v in tile_elem.items():\n    print(k, v.shape)\ntile_elem['edges_adjecency']","metadata":{"_uuid":"18913258-fcb0-4042-86ae-aac44f3003c5","_cell_guid":"14bb7a7f-711f-400d-a59c-7954eea7b48a","scrolled":true,"execution":{"iopub.status.busy":"2023-11-20T03:54:29.244819Z","iopub.execute_input":"2023-11-20T03:54:29.245265Z","iopub.status.idle":"2023-11-20T03:54:29.276558Z","shell.execute_reply.started":"2023-11-20T03:54:29.245224Z","shell.execute_reply":"2023-11-20T03:54:29.275442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Collator \n\nThis collator collects and organise tile dataset","metadata":{"_uuid":"8ca2a236-3705-4e93-838c-b38ef91debbc","_cell_guid":"56579611-c2ca-4b25-85da-75fe7a70d4d8","trusted":true}},{"cell_type":"code","source":"@dataclass\nclass GraphCollator:\n    pad_to_multiple_of: int = 64\n    targets:bool = True\n    padding_idx:int = 120\n    node_padding_idx:int = 0\n        \n        \n    def pad_edge_adjacency(self, 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    \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, padding_value=self.padding_idx),\n                                      (0, node_pad_amount), value=self.padding_idx).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(self.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\n        max_node_config_len = max([elem['node_config_ids'].shape[0] for elem in batch])\n        node_config_pad_amount = self.pad_to_multiple_of - max_node_config_len % max(self.pad_to_multiple_of, 1)\n        output['node_config_ids'] = F.pad(pad_sequence([elem['node_config_ids'] for elem in batch], batch_first=True),\n                                         (0, node_config_pad_amount), value=0).long()\n        padded_node_config_feat = pad_sequence([elem['node_config_feat'].permute(1,0,2) for elem in batch], batch_first=True, padding_value=-1)\n        padded_node_config_feat = F.pad(padded_node_config_feat.permute(0,2,1,3),\n                                           (0,0,0, node_config_pad_amount,0,0), value=-1)\n        \n        output['node_config_feat'] = torch.where(padded_node_config_feat!=-1, padded_node_config_feat, self.node_padding_idx)\n                                      \n        output['config_idxs'] = torch.stack([torch.from_numpy(elem['selected_idxs']) for elem in batch])\n        \n        if self.targets:\n            output['config_runtime'] = pad_sequence([elem['config_runtime'].float() for elem in batch], batch_first=True)\n        return output","metadata":{"_uuid":"5aaae820-be59-44b5-97a4-33d3729bfee8","_cell_guid":"e8f940ec-77c0-4d43-9825-964671417596","execution":{"iopub.status.busy":"2023-11-20T03:54:29.277947Z","iopub.execute_input":"2023-11-20T03:54:29.278498Z","iopub.status.idle":"2023-11-20T03:54:29.297490Z","shell.execute_reply.started":"2023-11-20T03:54:29.278455Z","shell.execute_reply":"2023-11-20T03:54:29.296313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"collate_fn = GraphCollator(64)\nbatch = collate_fn([tile_dataset[0], tile_dataset[1]])\nfor k,v in batch.items():\n    print(k,v.shape)","metadata":{"_uuid":"f95e3b5d-5402-4ad1-b5f8-632f2fe7ba8d","_cell_guid":"9f30c6e9-e937-46c4-84ec-b4adb0fc2d92","execution":{"iopub.status.busy":"2023-11-20T03:54:29.299094Z","iopub.execute_input":"2023-11-20T03:54:29.300439Z","iopub.status.idle":"2023-11-20T03:54:29.327172Z","shell.execute_reply.started":"2023-11-20T03:54:29.300371Z","shell.execute_reply":"2023-11-20T03:54:29.326117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Bert-like Tile Model","metadata":{"_uuid":"a3f588e7-bc4e-4047-8342-531a97d34a88","_cell_guid":"d04b18cc-190b-41d1-abe4-1c0ee08ef6c5","trusted":true}},{"cell_type":"code","source":"@dataclass\nclass GraphConfig:\n    num_hidden_layers: int = 10\n    hidden_size: int = 256\n    num_attention_heads: int = 16\n    intermediate_size: int = 64\n    chunk_size_feed_forward: int = 64\n    attention_probs_dropout_prob: float = 0.0\n    max_position_embeddings: int = 512\n    hidden_dropout_prob: float = 0.0\n    layer_norm_eps: float = 1e-11\n    hidden_act: str = 'gelu'\n    initializer_range: float = 0.0135\n    output_hidden_states: bool = False\n    output_attentions: bool = False\n    gradient_checkpointing: bool = False\n    margin: float = 0.035\n    number_permutations: int = 15\n    \n    def __post_init__(self):\n        self.embedding_size = self.hidden_size\n    \n    def validate(self):\n        if self.hidden_size % self.num_attention_heads != 0 and not hasattr(self, \"embedding_size\"):\n            raise ValueError(\n                f\"The hidden size ({self.hidden_size}) is not a multiple of the number of attention \"\n                f\"heads ({self.num_attention_heads})\"\n            )\n            \n    def save_config(self, path):\n        config = asdict(self)\n        with open(path, 'w') as f:\n            json.dump(config, f)\n            \n    @classmethod\n    def load_config(cls, path):\n        with open(path, 'r') as f:\n            config = json.load(f)\n        return cls(**config)","metadata":{"_uuid":"73dce51d-da61-4118-afa3-3d091ccf96e6","_cell_guid":"46de7eed-05bd-48da-938c-4b2b2795bd7c","execution":{"iopub.status.busy":"2023-11-20T03:54:29.328351Z","iopub.execute_input":"2023-11-20T03:54:29.328718Z","iopub.status.idle":"2023-11-20T03:54:29.342766Z","shell.execute_reply.started":"2023-11-20T03:54:29.328688Z","shell.execute_reply":"2023-11-20T03:54:29.341547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loss \nClass `MultiElementRankLoss`\n* Uses Ranking loss to compare different configuration\n* Compares does configurations with different indexes, masks those cases where the permutation returns the same element\n* Compares multiple configurations in each run","metadata":{"_uuid":"3795f832-f389-420c-af43-549c8c2b4173","_cell_guid":"36269ea9-434d-4852-9c53-44e0d6c636a3","trusted":true}},{"cell_type":"code","source":"class MultiElementRankLoss(nn.Module):\n    \"\"\"\n    Loss function that compares the output of the model with the output of the model with a permutation of the elements\n    \"\"\"\n    \n    def __init__(self, margin:float=0.0, number_permutations:int = 1) -> None:\n        super().__init__()\n        self.loss_fn = torch.nn.MarginRankingLoss(margin=margin, reduction = 'none')\n        self.number_permutations = number_permutations\n    \n    def calculate_rank_loss(self,\n                            outputs: torch.Tensor,\n                            config_runtime: torch.Tensor,\n                            config_idxs: torch.Tensor\n                            ):\n        \"\"\"\n        Generates a permutation of the predictions and targets and calculates the loss MarginRankingLoss against the permutation\n        Args:\n            outputs: Tensor of shape (bs, seq_len) with the outputs of the model\n            config_runtime: Tensor of shape (bs, seq_len) with the runtime of the model\n            config_mask: Tensor of shape (bs, seq_len) with 1 in the positions of the elements\n            and 0 in the positions of the padding\n        Returns:\n            loss: Tensor of shape (bs, seq_len) with the loss for each element in the batch\n        \"\"\"\n        bs, num_configs = outputs.shape\n        permutation = torch.randperm(num_configs) \n        permuted_idxs = config_idxs[:, permutation]\n        # We mask those cases where we compare the same configuration\n        config_mask = torch.where(config_idxs != permuted_idxs, 1, 0)\n        permuted_runtime = config_runtime[:, permutation]\n        labels = 2*((config_runtime - permuted_runtime) > 0) -1\n        permuted_output = outputs[:, permutation]\n        loss = self.loss_fn(outputs.view(-1,1), permuted_output.view(-1,1), labels.view(-1,1))\n        loss = loss.view(bs, num_configs) * config_mask\n        return loss.mean()\n                \n    \n    def forward(self, outputs: torch.Tensor, config_runtime: torch.Tensor, config_idxs: torch.Tensor):\n        loss = 0 \n        for _ in range(self.number_permutations):\n            loss += self.calculate_rank_loss(outputs, config_runtime, config_idxs)\n        return loss / self.number_permutations","metadata":{"_uuid":"5c38d8b0-9273-45ac-addc-3bab9e19e9b6","_cell_guid":"0ca046a3-960a-4de7-8c00-da9880288d49","execution":{"iopub.status.busy":"2023-11-20T03:54:29.344544Z","iopub.execute_input":"2023-11-20T03:54:29.344970Z","iopub.status.idle":"2023-11-20T03:54:29.359261Z","shell.execute_reply.started":"2023-11-20T03:54:29.344933Z","shell.execute_reply":"2023-11-20T03:54:29.358131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Metric\nClass `TileTopK` Compute the evaluation metric of predict runtimes and target runtims of top K observations","metadata":{"_uuid":"b0759b49-37e8-46ba-a70a-3c7124bb10ee","_cell_guid":"78958029-5404-4b32-b363-5389238137ba","trusted":true}},{"cell_type":"code","source":"class TileTopK(tm.Metric):\n    higher_is_better = True\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        \"\"\"\n        Update the metric state\n        Args:\n            preds: Tensor of shape (bs, seq_len) with the predicted runtimes orders\n            target: Tensor of shape (bs, seq_len) with the target runtimes\n            config_attn_mask: Tensor of shape (bs, seq_len) with 1 in the positions of the elements\n        \"\"\"\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(config_attn_mask.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()","metadata":{"_uuid":"2aaf2572-24ab-445d-a4b0-e473d45252c4","_cell_guid":"16b23b0d-94f1-43f9-8e62-ee6da8d23e69","execution":{"iopub.status.busy":"2023-11-20T03:54:29.360720Z","iopub.execute_input":"2023-11-20T03:54:29.361034Z","iopub.status.idle":"2023-11-20T03:54:29.375545Z","shell.execute_reply.started":"2023-11-20T03:54:29.361008Z","shell.execute_reply":"2023-11-20T03:54:29.374328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## BERT like (GATs) Model\nThis Bert implementation is based on [Graph Attention Networks](https://arxiv.org/abs/1710.10903)\n* Removed the parts corresponding to Cross-attention\n* Made layer_head_mask the same for all layers, heads\n* The Head mask corresponds to the edge adjacency\n\nThis section includes BERTEncoder, BertLayer, BertIntermediate, BertOutput, BertAttention, ... classes \nThese classes are modified from https://github.com/huggingface/transformers/blob/main/src/transformers/models/bert/modeling_bert.py","metadata":{"_uuid":"f339c0f7-1664-4ccd-a07d-425ed9e077ca","_cell_guid":"d39b91b5-51be-4195-b16d-ccfa2e0c4d79","trusted":true}},{"cell_type":"code","source":"class BertEncoder(nn.Module):\n    def __init__(self, config:GraphConfig):\n        super().__init__()\n        self.config = config\n        self.layer = nn.ModuleList([BertLayer(config) for _ in range(config.num_hidden_layers)])\n        self.gradient_checkpointing = False\n\n    def forward(self,\n                hidden_states: torch.Tensor,\n                attention_mask: Optional[torch.FloatTensor] = None,\n                head_mask: Optional[torch.FloatTensor] = None,\n                output_attentions: Optional[bool] = False,\n                output_hidden_states: Optional[bool] = False,\n                return_dict: Optional[bool] = True,\n                ) -> Union[Tuple[torch.Tensor], BaseModelOutputWithPastAndCrossAttentions]:\n        all_hidden_states = () if output_hidden_states else None\n        all_self_attentions = () if output_attentions else None\n\n        for i, layer_module in enumerate(self.layer):\n            if output_hidden_states:\n                all_hidden_states = all_hidden_states + (hidden_states,)\n\n            layer_head_mask = head_mask #DONE: Same Head Mask for all layers\n\n            if self.gradient_checkpointing and self.training:\n\n                def create_custom_forward(module):\n                    def custom_forward(*inputs):\n                        return module(*inputs,  output_attentions)\n\n                    return custom_forward\n\n                layer_outputs = torch.utils.checkpoint.checkpoint(\n                    create_custom_forward(layer_module),\n                    hidden_states,\n                    attention_mask,\n                    layer_head_mask,\n                )\n            else:\n                layer_outputs = layer_module(\n                    hidden_states,\n                    attention_mask,\n                    layer_head_mask,\n                    output_attentions,\n                )\n\n            hidden_states = layer_outputs[0]\n            if output_attentions:\n                all_self_attentions = all_self_attentions + (layer_outputs[1],)\n\n        if output_hidden_states:\n            all_hidden_states = all_hidden_states + (hidden_states,)\n\n        if not return_dict:\n            return tuple(\n                v\n                for v in [\n                    hidden_states,\n                    all_hidden_states,\n                    all_self_attentions,\n                ]\n                if v is not None\n            )\n        return BaseModelOutputWithPastAndCrossAttentions(\n            last_hidden_state=hidden_states,\n            past_key_values=None,\n            hidden_states=all_hidden_states,\n            attentions=all_self_attentions,\n            cross_attentions=None,\n        )\n        \n        \nclass BertLayer(nn.Module):\n    def __init__(self, config:GraphConfig):\n        super().__init__()\n        self.chunk_size_feed_forward = config.chunk_size_feed_forward\n        self.seq_len_dim = 1\n        self.attention = BertAttention(config)\n        self.intermediate = BertIntermediate(config)\n        self.output = BertOutput(config)\n\n    def forward(\n        self,\n        hidden_states: torch.Tensor,\n        attention_mask: Optional[torch.FloatTensor] = None,\n        head_mask: Optional[torch.FloatTensor] = None,\n        output_attentions: Optional[bool] = False,\n    ) -> Tuple[torch.Tensor]:\n        # decoder uni-directional self-attention cached key/values tuple is at positions 1,2\n        self_attention_outputs = self.attention(\n            hidden_states,\n            attention_mask,\n            head_mask,\n            output_attentions=output_attentions,\n        )\n        attention_output = self_attention_outputs[0]\n        outputs = self_attention_outputs[1:]  # add self attentions if we output attention weights\n        layer_output = apply_chunking_to_forward(\n            self.feed_forward_chunk, self.chunk_size_feed_forward, self.seq_len_dim, attention_output\n        )\n        outputs = (layer_output,) + outputs\n\n\n        return outputs\n\n    def feed_forward_chunk(self, attention_output):\n        intermediate_output = self.intermediate(attention_output)\n        layer_output = self.output(intermediate_output, attention_output)\n        return layer_output\n    \nclass BertIntermediate(nn.Module):\n    def __init__(self, config:GraphConfig):\n        super().__init__()\n        self.dense = nn.Linear(config.hidden_size, config.intermediate_size)\n        if isinstance(config.hidden_act, str):\n            self.intermediate_act_fn = ACT2FN[config.hidden_act]\n        else:\n            self.intermediate_act_fn = config.hidden_act\n\n    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:\n        hidden_states = self.dense(hidden_states)\n        hidden_states = self.intermediate_act_fn(hidden_states)\n        return hidden_states\n    \nclass BertOutput(nn.Module):\n    def __init__(self, config:GraphConfig):\n        super().__init__()\n        self.dense = nn.Linear(config.intermediate_size, config.hidden_size)\n        self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)\n        self.dropout = nn.Dropout(config.hidden_dropout_prob)\n\n    def forward(self, hidden_states: torch.Tensor, input_tensor: torch.Tensor) -> torch.Tensor:\n        hidden_states = self.dense(hidden_states)\n        hidden_states = self.dropout(hidden_states)\n        hidden_states = self.LayerNorm(hidden_states + input_tensor)\n        return hidden_states\n    \nclass BertAttention(nn.Module):\n    def __init__(self, config:GraphConfig, position_embedding_type=None):\n        super().__init__()\n        self.self = BertSelfAttention(config, position_embedding_type=position_embedding_type)\n        self.output = BertSelfOutput(config)\n        self.pruned_heads = set()\n\n    def forward(\n        self,\n        hidden_states: torch.Tensor,\n        attention_mask: Optional[torch.FloatTensor] = None,\n        head_mask: Optional[torch.FloatTensor] = None,\n        output_attentions: Optional[bool] = False,\n    ) -> Tuple[torch.Tensor]:\n        self_outputs = self.self(\n            hidden_states,\n            attention_mask,\n            head_mask,\n            output_attentions,\n        )\n        attention_output = self.output(self_outputs[0], hidden_states)\n        outputs = (attention_output,) + self_outputs[1:]  # add attentions if we output them\n        return outputs\n    \n    \nclass BertSelfAttention(nn.Module):\n    def __init__(self, config:GraphConfig, position_embedding_type=None):\n        super().__init__()\n        if config.hidden_size % config.num_attention_heads != 0 and not hasattr(config, \"embedding_size\"):\n            raise ValueError(\n                f\"The hidden size ({config.hidden_size}) is not a multiple of the number of attention \"\n                f\"heads ({config.num_attention_heads})\"\n            )\n\n        self.num_attention_heads = config.num_attention_heads\n        self.attention_head_size = int(config.hidden_size / config.num_attention_heads)\n        self.all_head_size = self.num_attention_heads * self.attention_head_size\n\n        self.query = nn.Linear(config.hidden_size, self.all_head_size)\n        self.key = nn.Linear(config.hidden_size, self.all_head_size)\n        self.value = nn.Linear(config.hidden_size, self.all_head_size)\n\n        self.dropout = nn.Dropout(config.attention_probs_dropout_prob)\n        self.position_embedding_type = position_embedding_type or getattr(\n            config, \"position_embedding_type\", \"absolute\"\n        )\n        if self.position_embedding_type == \"relative_key\" or self.position_embedding_type == \"relative_key_query\":\n            self.max_position_embeddings = config.max_position_embeddings\n            self.distance_embedding = nn.Embedding(2 * config.max_position_embeddings - 1, self.attention_head_size)\n\n\n    def transpose_for_scores(self, x: torch.Tensor) -> torch.Tensor:\n        new_x_shape = x.size()[:-1] + (self.num_attention_heads, self.attention_head_size)\n        x = x.view(new_x_shape)\n        return x.permute(0, 2, 1, 3)\n\n    def forward(\n        self,\n        hidden_states: torch.Tensor,\n        attention_mask: Optional[torch.FloatTensor] = None,\n        head_mask: Optional[torch.FloatTensor] = None,\n        output_attentions: Optional[bool] = False,\n    ) -> Tuple[torch.Tensor]:\n        \n        mixed_query_layer = self.query(hidden_states)\n        key_layer = self.transpose_for_scores(self.key(hidden_states))\n        value_layer = self.transpose_for_scores(self.value(hidden_states))\n        query_layer = self.transpose_for_scores(mixed_query_layer)\n\n\n        # Take the dot product between \"query\" and \"key\" to get the raw attention scores.\n        attention_scores = torch.matmul(query_layer, key_layer.transpose(-1, -2))\n\n        if self.position_embedding_type == \"relative_key\" or self.position_embedding_type == \"relative_key_query\":\n            query_length, key_length = query_layer.shape[2], key_layer.shape[2]\n            position_ids_l = torch.arange(query_length, dtype=torch.long, device=hidden_states.device).view(-1, 1)\n            position_ids_r = torch.arange(key_length, dtype=torch.long, device=hidden_states.device).view(1, -1)\n            distance = position_ids_l - position_ids_r\n\n            positional_embedding = self.distance_embedding(distance + self.max_position_embeddings - 1)\n            positional_embedding = positional_embedding.to(dtype=query_layer.dtype)  # fp16 compatibility\n\n            if self.position_embedding_type == \"relative_key\":\n                relative_position_scores = torch.einsum(\"bhld,lrd->bhlr\", query_layer, positional_embedding)\n                attention_scores = attention_scores + relative_position_scores\n            elif self.position_embedding_type == \"relative_key_query\":\n                relative_position_scores_query = torch.einsum(\"bhld,lrd->bhlr\", query_layer, positional_embedding)\n                relative_position_scores_key = torch.einsum(\"bhrd,lrd->bhlr\", key_layer, positional_embedding)\n                attention_scores = attention_scores + relative_position_scores_query + relative_position_scores_key\n\n        attention_scores = attention_scores / math.sqrt(self.attention_head_size)\n        if attention_mask is not None:\n            # Apply the attention mask is (precomputed for all layers in BertModel forward() function)\n            attention_scores = attention_scores + attention_mask\n\n        # Normalize the attention scores to probabilities.\n        attention_probs = nn.functional.softmax(attention_scores, dim=-1)\n\n        # This is actually dropping out entire tokens to attend to, which might\n        # seem a bit unusual, but is taken from the original Transformer paper.\n        attention_probs = self.dropout(attention_probs)\n\n        # Mask heads if we want to\n        if head_mask is not None:\n            attention_probs = attention_probs * head_mask #DONE: Same Head Mask for all Heads\n\n        context_layer = torch.matmul(attention_probs, value_layer)\n\n        context_layer = context_layer.permute(0, 2, 1, 3).contiguous()\n        new_context_layer_shape = context_layer.size()[:-2] + (self.all_head_size,)\n        context_layer = context_layer.view(new_context_layer_shape)\n\n        outputs = (context_layer, attention_probs) if output_attentions else (context_layer,)\n\n        return outputs\n\n\nclass BertSelfOutput(nn.Module):\n    def __init__(self, config:GraphConfig):\n        super().__init__()\n        self.dense = nn.Linear(config.hidden_size, config.hidden_size)\n        self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)\n        self.dropout = nn.Dropout(config.hidden_dropout_prob)\n\n    def forward(self, hidden_states: torch.Tensor, input_tensor: torch.Tensor) -> torch.Tensor:\n        hidden_states = self.dense(hidden_states)\n        hidden_states = self.dropout(hidden_states)\n        hidden_states = self.LayerNorm(hidden_states + input_tensor)\n        return hidden_states\n    \n    \nclass NodeEncoder(nn.Module):\n    \n    def __init__(self, config:GraphConfig):\n        super().__init__()\n        self.node_opcode_embeddings = nn.Embedding(NODE_OP_CODES+1 , config.embedding_size, padding_idx=NODE_OP_CODES)\n        self.linear = nn.Linear(NODE_FEATS, config.embedding_size, bias=False)\n        self.layer_norm = nn.LayerNorm(config.embedding_size, eps=config.layer_norm_eps)\n        \n        \n    def forward(self,\n                node_opcode: torch.Tensor,\n                node_feat: torch.Tensor\n                ) -> torch.Tensor:\n        opcode_embeddings = self.node_opcode_embeddings(node_opcode) \n        node_feats =  self.linear(node_feat)\n        features = opcode_embeddings + node_feats\n        features = self.layer_norm(features)\n        return features\n    \n    \nclass BertNodeEncoder(nn.Module):\n    \n    def __init__(self, config:GraphConfig) -> None:\n        super().__init__()\n        self.config = config\n        self.node_embeddings = NodeEncoder(config)\n        self.node_encoder = BertEncoder(config)\n        \n    def forward(self,\n                node_opcode: torch.Tensor,\n                node_feat: torch.Tensor,\n                edges_adjecency: torch.Tensor,\n                node_attn_mask: torch.Tensor\n                ):\n        node_embeddings = self.node_embeddings(node_opcode, node_feat)\n        node_attn_mask = node_attn_mask.unsqueeze(1).unsqueeze(-1)\n        node_encoder_outputs = self.node_encoder(node_embeddings,\n                                                 attention_mask=node_attn_mask,\n                                                 head_mask=edges_adjecency.unsqueeze(0).repeat(self.config.num_hidden_layers, 1, 1, 1).unsqueeze(2),\n                                                 output_attentions=True)\n        return node_encoder_outputs\n    \ndef transform_node_positional_embeddings(embeddings_output:torch.Tensor,\n                                         node_config_ids:torch.Tensor,\n                                         num_nodes:int\n                                         ) -> torch.Tensor:\n    bs, num_configs, _, dim = embeddings_output.shape\n    idxs = node_config_ids.unsqueeze(1).repeat(1,num_configs,1)\n    zeros = torch.zeros(bs, num_configs, num_nodes, dim, device=embeddings_output.device, dtype=embeddings_output.dtype)\n    idxs = idxs.unsqueeze(-1).repeat(1,1,1,dim)\n    zeros.scatter_reduce_(2, idxs, embeddings_output, reduce='sum')\n    return zeros\n\nclass NodeFeatEmbeddings(nn.Module):\n    def __init__(self, config:GraphConfig):\n        super().__init__()\n        self.config = config\n        self.node_feat_embeddings = nn.Linear(NODE_CONFIG_FEATS + CONFIG_FEATS, config.embedding_size, bias=False)\n        self.layer_norm = nn.LayerNorm(config.embedding_size, eps=config.layer_norm_eps)\n        \n    def forward(self, node_config_feat: torch.Tensor, node_config_ids: torch.Tensor, num_nodes:int) -> torch.Tensor:\n        node_config_feat_embeddings = self.node_feat_embeddings(node_config_feat)\n        node_config_feat_embeddings = self.layer_norm(node_config_feat_embeddings)\n        node_config_feat_embeddings = transform_node_positional_embeddings(node_config_feat_embeddings, node_config_ids, num_nodes)\n        return node_config_feat_embeddings\n        \n    \nclass BertGraphEncoder(nn.Module):\n    def __init__(self, config:GraphConfig) -> None:\n        super().__init__()\n        self.config = config\n        self.node_embeddings = NodeEncoder(config)\n        self.node_encoder = BertEncoder(config)\n        self.node_feat_embeddings = NodeFeatEmbeddings(config)\n        \n    def forward(self,\n                node_opcode: torch.Tensor, # (bs, num_nodes)\n                node_feat: torch.Tensor, # (bs, num_nodes, num_node_feats)\n                edges_adjecency: torch.Tensor, # (bs, num_nodes, num_nodes)\n                node_attn_mask: torch.Tensor, # (bs, num_nodes)\n                node_config_feat: torch.Tensor, # (bs, num_configs, num_config_nodes, num_node_feats)\n                node_config_ids: torch.Tensor, # (bs, num_configs, num_config_nodes)\n                ):\n        bs, num_nodes = node_opcode.shape\n        num_configs = node_config_feat.shape[1]\n        node_embeddings = self.node_embeddings(node_opcode, node_feat)\n        node_config_feat_embeddings = self.node_feat_embeddings(node_config_feat, node_config_ids, num_nodes)\n        \n        node_embeddings = node_embeddings.unsqueeze(1).repeat(1, num_configs, 1, 1)\n        node_embeddings += node_config_feat_embeddings\n        node_attn_mask = node_attn_mask.unsqueeze(1).repeat(1, num_configs, 1)\n        node_embeddings = node_embeddings.reshape(bs *num_configs, num_nodes, -1)\n        node_attn_mask = node_attn_mask.reshape(bs *num_configs, num_nodes)\n        node_attn_mask = node_attn_mask.unsqueeze(1).unsqueeze(-1)\n        edges_adjecency = edges_adjecency.unsqueeze(1).repeat(1, num_configs, 1, 1).reshape(bs *num_configs, num_nodes, num_nodes)\n        edges_adjecency = edges_adjecency.unsqueeze(1)\n        \n\n        node_encoder_outputs = self.node_encoder(node_embeddings,\n                                                 attention_mask=node_attn_mask,\n                                                 head_mask=edges_adjecency,\n                                                 output_attentions=True)\n        \n        return node_encoder_outputs.last_hidden_state.reshape(bs, num_configs, num_nodes, -1)","metadata":{"_uuid":"91ae303f-505b-4be4-a563-6b76fd9ea075","_cell_guid":"83e94f1c-ea80-459f-a737-6e20d98299b4","execution":{"iopub.status.busy":"2023-11-20T03:54:29.377568Z","iopub.execute_input":"2023-11-20T03:54:29.377972Z","iopub.status.idle":"2023-11-20T03:54:29.459141Z","shell.execute_reply.started":"2023-11-20T03:54:29.377934Z","shell.execute_reply":"2023-11-20T03:54:29.457972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## BertGraph Encoder\n`GraphEncoder`","metadata":{"_uuid":"f76f68b1-f638-45f2-90a4-a3a8ff532262","_cell_guid":"ce48a35f-29f4-41ca-8a02-95e7b88b7bf5","trusted":true}},{"cell_type":"code","source":"class GraphEncoder(nn.Module):\n    \n    config_class = GraphConfig\n    \n    def __init__(self, config:GraphConfig):\n        super().__init__()\n        self.config = config\n        self.node_encoder = BertGraphEncoder(config)\n        self.head = nn.Linear(config.hidden_size, 1)\n        self.loss_fn = MultiElementRankLoss(margin=config.margin,\n                                            number_permutations=config.number_permutations)\n        \n    def forward(self,\n                node_opcode: torch.Tensor, # (bs, num_nodes)\n                node_feat: torch.Tensor, # (bs, num_nodes, num_node_feats)\n                edges_adjecency: torch.Tensor, # (bs, num_nodes, num_nodes)\n                node_attn_mask: torch.Tensor, # (bs, num_nodes)\n                node_config_feat: torch.Tensor, # (bs, num_configs, num_config_nodes, num_node_feats)\n                node_config_ids: torch.Tensor, # (bs, num_configs, num_config_nodes)\n                config_idxs: Optional[torch.Tensor] = None, # (bs, num_configs)\n                config_runtime: Optional[torch.Tensor] = None,):\n        \n        last_hidden_state = self.node_encoder(node_opcode,\n                                    node_feat,\n                                    edges_adjecency,\n                                    node_attn_mask,\n                                    node_config_feat,\n                                    node_config_ids)\n        \n        output = self.head(last_hidden_state[:,:,0]).squeeze(-1)\n        outputs = {'outputs': output, 'order': torch.argsort(output, dim=1)}\n        if config_runtime is not None:\n            loss = 0\n            loss += self.loss_fn(output, config_runtime, config_idxs)\n            outputs['loss'] = loss\n        return outputs","metadata":{"_uuid":"f13c6823-4221-476e-a7b5-30180937e203","_cell_guid":"ad73a9a2-4661-4273-982a-086be188dbc6","execution":{"iopub.status.busy":"2023-11-20T03:54:29.462350Z","iopub.execute_input":"2023-11-20T03:54:29.463460Z","iopub.status.idle":"2023-11-20T03:54:29.477284Z","shell.execute_reply.started":"2023-11-20T03:54:29.463414Z","shell.execute_reply":"2023-11-20T03:54:29.475845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## A wrapper for Lighting AI\nWe use lighting library to optimize the model","metadata":{"_uuid":"0b8a306a-4374-4802-9678-1051d5ca5687","_cell_guid":"e8642954-71cb-404c-bf9b-c6299e1ca1fc","trusted":true}},{"cell_type":"code","source":"from transformers import get_linear_schedule_with_warmup, get_cosine_schedule_with_warmup\nfrom pytorch_lightning.callbacks.early_stopping import EarlyStopping\n\nclass LightningWrapper(pl.LightningModule):\n    def __init__(self, model, fold, params):\n        super().__init__()\n        self.model = model\n        self.topk = TileTopK()\n        self.fold = fold\n        self.model_path = f'/kaggle/tmp/fold{fold}/BertGraph_fold{fold}_best.pth'\n        self.params = params\n        self.best_score = 1.0 # best score\n        self.train_loss = []\n        self.val_loss = []\n        self._step = 0\n        \n    def forward(self, x):\n        return self.model(x)\n    \n    def training_step(self, batch, batch_idx):\n        for k, v in batch.items():\n            batch[k] = v.to(device)\n        batch_size = len(batch.items())\n        \n        outputs = self.model(**batch)\n        loss = outputs['loss']\n        y_preds = outputs['outputs']\n        y_order = outputs['order']\n        self.train_loss.extend([loss.item()]* batch_size)\n        return loss\n    \n    def on_train_epoch_end(self):\n        avg_loss = np.mean(self.train_loss)\n        print(f'Loss: {avg_loss:.4f} '\n              f'LR: {self.params[\"lr\"]:.8f}  ')\n        self.train_loss.clear()  # free memory\n        self._step = 0\n\n    def validation_step(self, batch, batch_idx):\n        for k, v in batch.items():\n            batch[k] = v.to(device)\n        batch_size = len(batch.items())\n        with torch.cuda.amp.autocast(enabled=self.params['apex']):\n            outputs = self.model(**batch)\n        loss = outputs['loss']\n        self.log(\"val_loss\", loss, prog_bar=True)\n        self.val_loss.extend([loss.item()]* batch_size)\n        \n        config_attn_mask = torch.ones_like(batch['config_runtime'], device=batch['config_runtime'].device)\n        self.topk.update(outputs['outputs'], batch['config_runtime'], config_attn_mask)\n        return loss\n    \n    def on_validation_end(self) -> None:\n        topk = self.topk.compute().cpu()\n        self.print(f\"topk {topk:.3f}\")\n        avg_loss = np.mean(self.val_loss)\n        if avg_loss < self.best_score: \n            self.best_score = avg_loss\n            print(f\"Found a better model for fold {self.fold} with score = {self.best_score: .3f}\\n\"\n                  f\"Save it to {self.model_path}\")\n            # Save the model\n            torch.save(self.model.state_dict(), self.model_path)  \n        \n        self.topk.reset()\n        self.val_loss.clear()\n        return super().on_validation_end()\n\n    def test_step(self, batch, batch_idx):\n        x, y = batch\n        y_hat = self.model(x)\n        loss = self.model.loss(y_hat, y)\n        self.log(\"test_loss\", loss, prog_bar=True)\n        return loss\n    \n    def configure_optimizers(self):\n        self.optimizer = torch.optim.AdamW(self.trainer.model.parameters(), lr=self.params['lr'],\n                                      eps=self.params['eps'], betas=self.params['betas'],\n                                      maximize=False)\n        return self.optimizer","metadata":{"_uuid":"85d65927-f169-435a-b624-5d7f87d40cae","_cell_guid":"89969abe-686b-4994-b445-a583fda24e29","execution":{"iopub.status.busy":"2023-11-20T03:54:29.479099Z","iopub.execute_input":"2023-11-20T03:54:29.479579Z","iopub.status.idle":"2023-11-20T03:54:29.506447Z","shell.execute_reply.started":"2023-11-20T03:54:29.479541Z","shell.execute_reply":"2023-11-20T03:54:29.505018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Cross Validation Training for tile model","metadata":{"_uuid":"7f98baa5-f3a6-44b4-a8aa-918f715ad05e","_cell_guid":"4867e97b-aa70-4ee6-a02d-cca23dd97aaa","trusted":true}},{"cell_type":"markdown","source":"## CV Splits","metadata":{"_uuid":"4d1fccff-23ae-44f4-b1d7-1647b83d35e3","_cell_guid":"4fbaca57-aa38-4463-9381-108e5ca66e22","trusted":true}},{"cell_type":"code","source":"from sklearn.model_selection import StratifiedKFold\nn_fold = 5\nseed = 42\n# Create CV splits for tile df\ntile_train_df = tile_df[tile_df['split'] != 'test'].reset_index(drop=True)\ntile_test_df = tile_df[tile_df['split'] == 'test'].reset_index(drop=True)\nSKFold = StratifiedKFold(n_splits=n_fold, shuffle=True, random_state=seed)\nfor n, (train_index, val_index) in enumerate(SKFold.split(tile_train_df, tile_train_df[['collection']])):\n    tile_train_df.loc[val_index, 'fold'] = int(n)\ntile_train_df['fold'] = tile_train_df['fold'].astype(int)\n# print(tile_train_df['split'].value_counts())\n# display(tile_train_df.groupby('fold').size())\ndisplay(tile_train_df.head(3))","metadata":{"_uuid":"d68299f0-a75e-49db-bf5e-73d85d95f4d9","_cell_guid":"e86b4feb-d2d1-4b44-90b4-f6d535fe1086","execution":{"iopub.status.busy":"2023-11-20T03:54:29.508247Z","iopub.execute_input":"2023-11-20T03:54:29.508747Z","iopub.status.idle":"2023-11-20T03:54:29.566598Z","shell.execute_reply.started":"2023-11-20T03:54:29.508700Z","shell.execute_reply":"2023-11-20T03:54:29.565719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## CV Training with lightning AI","metadata":{}},{"cell_type":"code","source":"import shutil\nfrom pathlib import Path\n\n# torch.set_float32_matmul_precision(\"medium\") # Set \ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\n\nOUTPUT_DIR = \"/kaggle/working/\"\nconfig_kwargs = {'hidden_size': 128,\n                 'num_attention_heads': 4,\n                 'num_hidden_layers': 2,\n                 'intermediate_size': 64,\n                 'gradient_checkpointing': True,\n                 'margin': 0.1,\n                 'number_permutations': 6}\n\nconfig = GraphConfig(**config_kwargs)\nencoder = GraphEncoder(config)","metadata":{"execution":{"iopub.status.busy":"2023-11-20T03:54:29.568014Z","iopub.execute_input":"2023-11-20T03:54:29.568648Z","iopub.status.idle":"2023-11-20T03:54:29.582776Z","shell.execute_reply.started":"2023-11-20T03:54:29.568607Z","shell.execute_reply":"2023-11-20T03:54:29.581450Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEV_ENABLE = True\nTOTAL_EPOCHS = 10\n# Training with lightning AI \ndef training_objective_pl(trial, mode, best_scores):\n    start = time.time()\n    print(f\"=== Begin the trial {trial.number}  ===\")\n    \n    params ={\n        \"batch_size\": 8,  # 8 for tile\n        \"lr\": trial.suggest_float(\"lr\", 1e-4, 1e-2),\n        \"gradient_accumulation_steps\": 2,\n        \"eps\": 1e-6,\n        \"betas\": (0.9, 0.999),\n        \"apex\": False, # if True, loss may explode\n        \"scheduler\": 'linear', # cosine/linear\n        \"epochs\": TOTAL_EPOCHS,\n        \"check_val_every_n_epoch\": 1,\n        \"max_grad_norm\": 1000,\n        \"seed\": 42,\n        \"batch_scheduler\": True,\n        \"num_cycles\": 0.5,\n        \"num_warmup_steps\": 50,\n    }\n    \n    # Create splits\n    total_scores = []\n    for idx, fold in enumerate(range(n_fold)):        \n        pl.seed_everything(42, workers=True) \n         \n        # Split the tile into train and valid datasets\n        train_df = tile_train_df[tile_train_df['fold']==fold].reset_index(drop=True)\n        valid_df = tile_train_df[tile_train_df['fold']!=fold].reset_index(drop=True)\n        train_dataset = TileDataset(train_df, num_configs=24)\n        valid_dataset = TileDataset(valid_df, num_configs=24)\n        \n        # Compute the training steps\n        params[\"num_train_steps\"] = int(len(train_df) / params['batch_size'] * params['epochs'])\n\n        # Create dataloaders\n        train_dataloader = DataLoader(train_dataset, collate_fn=collate_fn, batch_size=params['batch_size'],\n                                      shuffle=True, num_workers=4, pin_memory=True)\n        valid_dataloader = DataLoader(valid_dataset, collate_fn=collate_fn, batch_size=params['batch_size'],\n                                      num_workers=4, pin_memory=True)\n        # Load pretrained model\n        model_path = f\"/kaggle/input/bert-like-title-model-output/BertGraph_{mode}_fold{fold}_best.pth\"\n        encoder.load_state_dict(torch.load(model_path, map_location=torch.device('cpu')))\n        \n        \n        TMP_DIR = f'/kaggle/tmp/fold{fold}/' # Save training outputs   \n        Path(TMP_DIR).mkdir(parents=True, exist_ok=True)\n        # Early stop when the changes < 0.001\n        early_stop_callback = EarlyStopping(monitor=\"val_loss\", min_delta=0.00, patience=3,\n                                            verbose=True, mode=\"min\")\n        \n        # Training config\n        trainer_config = {\n            'default_root_dir': TMP_DIR,\n            'max_epochs': params['epochs'],\n            'precision': 32, # 'bf16-mixed',\n            'gradient_clip_val': 1.0,\n            'accumulate_grad_batches': 4,\n            'check_val_every_n_epoch': 2,\n            'accelerator': 'auto',\n            'fast_dev_run': DEV_ENABLE,  # True for developing\n            'enable_checkpointing': False, # False: not saving checkpoints\n            'callbacks': [early_stop_callback]\n        }\n        \n        trainer = pl.Trainer(**trainer_config)\n        # init the model directly on the device and with parameters in half-precision\n        # Create the model \n        model = LightningWrapper(model=encoder, fold=fold, params=params)\n        \n        trainer.fit(model, train_dataloader, valid_dataloader)\n        best_score = model.best_score # Use the average valid losses as the score\n        # Save the best model to working folder (/kaggle/working)\n        if best_score < best_scores[fold] :\n            best_scores[fold] = best_score\n            output_path = OUTPUT_DIR + f'BertGraph_{mode}_fold{fold}_best.pth'\n            # Overwrite the best model file\n            shutil.move(TMP_DIR+f'BertGraph_fold{fold}_best.pth', output_path)\n            print(f\"Current best score of fold {fold} is {best_scores[fold]:.3f}\\n\"\n                  f\"Save the model to {output_path}\")\n\n        total_scores.append(best_score)\n        del model\n        torch.cuda.empty_cache()\n        gc.collect()\n    \n    \n    avg_score = np.mean(total_scores) # Use the average valid loss as objective\n    print(f\"=== Finish the trial {trial.number} with {time.time() - start: .1f} seconds ===\")\n    return avg_score","metadata":{"execution":{"iopub.status.busy":"2023-11-20T03:54:29.584810Z","iopub.execute_input":"2023-11-20T03:54:29.585203Z","iopub.status.idle":"2023-11-20T03:54:29.605239Z","shell.execute_reply.started":"2023-11-20T03:54:29.585161Z","shell.execute_reply":"2023-11-20T03:54:29.604273Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"OUTPUT_DIR = \"/kaggle/working/\"\nIS_DEBUG = True\nmode = 'tile'\n \nN_TRIALS = 1\nTIME_OUT = 5000\nbest_scores = [1.0]*n_fold\n# training_objective_pl(best_scores)\nstudy = optuna.create_study(direction=\"minimize\")\n# # Training with lightning\nstudy.optimize(lambda trial: training_objective_pl(trial, mode, best_scores), n_trials=N_TRIALS, show_progress_bar=True)\nbest_params = study.best_trial.params\njoblib.dump(study, f\"{mode}_study.pkl\")\nprint(f\"Best model hyper value: {study.best_trial.value}\\nParameters: {best_params}\")","metadata":{"_uuid":"490c7a1e-d329-416d-b100-d034b76e43e6","_cell_guid":"d3115932-a5ca-4e46-bfe0-fa4ad0f00319","scrolled":true,"execution":{"iopub.status.busy":"2023-11-20T03:54:29.606841Z","iopub.execute_input":"2023-11-20T03:54:29.607239Z","iopub.status.idle":"2023-11-20T03:54:45.878649Z","shell.execute_reply.started":"2023-11-20T03:54:29.607208Z","shell.execute_reply":"2023-11-20T03:54:45.877344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{"_uuid":"2d941177-5e96-4ed2-a28e-cba16d3e6aa7","_cell_guid":"dbbb9013-3d83-4260-8e0a-6548b05bc900","trusted":true}},{"cell_type":"code","source":"def chunk_batch(batch, start_idx, end_idx):\n    output = {k:batch[k] for k in ['node_opcode', 'node_feat', 'edges_adjecency',\n                                   'node_attn_mask', 'node_config_ids']}\n    output['node_config_feat'] = batch['node_config_feat'][:, start_idx: end_idx]\n    return output\n\ndef generate_result_by_fold(fold):\n    split = 'test'\n    mode = 'tile'\n    print(f\"Starting infer the fold {fold}\")\n    \n    model = GraphEncoder(config)\n    model.to(device)\n    global dataset\n    collate_fn = GraphCollator(64, targets=split!=\"test\")\n    test_dataloader = DataLoader(dataset, batch_size=1, shuffle=False, num_workers=1,\n                                 collate_fn=collate_fn)\n    \n    path = OUTPUT_DIR+f\"BertGraph_{mode}_fold{fold}_best.pth\"\n    model.load_state_dict(torch.load(path, map_location=torch.device('cpu')))\n    model = model.eval()\n    predictions = [[] for i in range(len(dataset))] # Local predictions\n    for i, batch in enumerate(tqdm(test_dataloader)):\n        batch.pop('config_idxs')\n        batch = {k: v.to(device) for k, v in batch.items()}\n        # Chunk the configs to avoid OOM errors\n        num_configs = batch['node_config_feat'].shape[1]\n        configs_cut_points = list(range(0, num_configs, 100)) + [num_configs]\n        chunk_order = []\n        for start, end in zip(configs_cut_points, configs_cut_points[1:]):\n            chunked_batch = chunk_batch(batch, start, end)\n            with torch.no_grad():\n                output = model(**chunked_batch)\n            chunk_order.extend(output['outputs'].cpu().numpy())\n        predictions[i].append(np.concatenate(chunk_order))\n        \n    print(f\"Finish infering the fold {fold}\")\n    return predictions","metadata":{"execution":{"iopub.status.busy":"2023-11-20T03:54:45.880195Z","iopub.execute_input":"2023-11-20T03:54:45.880595Z","iopub.status.idle":"2023-11-20T03:54:45.896309Z","shell.execute_reply.started":"2023-11-20T03:54:45.880563Z","shell.execute_reply":"2023-11-20T03:54:45.895158Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import concurrent\nfrom concurrent.futures import ThreadPoolExecutor\nimport logging\nimport concurrent.futures\n# Create the model \nOUTPUT_DIR = \"/kaggle/working/\"\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\n\ndef make_tile_prediction():\n    tile_test_df = tile_df[tile_df['split'] == 'test'].reset_index(drop=True)\n    global dataset  # declare global variable\n    dataset = TileDataset(tile_test_df, num_configs=-1)\n    tile_xla_predictions = [[] for i in range(len(dataset))]\n    \n    # Use the best model in each fold to make the predictions\n    start = time.time()\n    print(f\"=== Begin to tile inference  ===\")\n    processes = []\n    with concurrent.futures.ThreadPoolExecutor(max_workers=3) as executor:\n        processes.append(executor.map(generate_result_by_fold, range(n_fold)))\n        for _ in concurrent.futures.as_completed(processes):\n            predictions = _.result()\n            print('predictions: ', predictions)            \n            # Update the prediction to global predictions\n            for i in range(len(predictions)):\n                tile_xla_predictions[i].extend(predictions[i])\n        \n#     for fold in range(n_fold):\n#         split = 'test'\n#         mode = 'tile'\n#         print(f\"Starting infer the fold {fold}\")\n\n#         model = GraphEncoder(config)\n#         model.to(device)\n#         collate_fn = GraphCollator(64, targets=split!=\"test\")\n#         test_dataloader = DataLoader(dataset, batch_size=1, shuffle=False, num_workers=1,\n#                                      collate_fn=collate_fn)\n\n#         path = OUTPUT_DIR+f\"BertGraph_{mode}_fold{fold}_best.pth\"\n#         print(path)\n#         if os.path.exists(path):\n#             model.load_state_dict(torch.load(path, map_location=torch.device('cpu')))\n#             model = model.eval()\n#             predictions = [[] for i in range(len(dataset))] # Local predictions\n#             for i, batch in enumerate(tqdm(test_dataloader)):\n#                 batch.pop('config_idxs')\n#                 batch = {k: v.to(device) for k, v in batch.items()}\n#                 # Chunk the configs to avoid OOM errors\n#                 num_configs = batch['node_config_feat'].shape[1]\n#                 configs_cut_points = list(range(0, num_configs, 100)) + [num_configs]\n#                 chunk_order = []\n#                 for start, end in zip(configs_cut_points, configs_cut_points[1:]):\n#                     chunked_batch = chunk_batch(batch, start, end)\n#                     with torch.no_grad():\n#                         output = model(**chunked_batch)\n#                     chunk_order.extend(output['outputs'].cpu().numpy())\n#                 predictions[i].append(np.concatenate(chunk_order))\n#             # Update the prediction to global predictions\n#             for i in range(len(dataset)):\n#                 tile_xla_predictions[i].extend(predictions[i])\n            \n    print(f\"=== Finish the tile inference with {time.time() - start: .1f} seconds ===\")\n        \n#     for fold in range(n_fold):\n#         generate_result_by_fold(fold)\n    print(tile_xla_predictions[0])\n    tile_xla_prediction_rankings = [np.argsort(np.mean(pred,axis=0))[:5] for pred in tile_xla_predictions] \n    idxs_string = [\";\".join(map(str,elem)) for elem in tile_xla_prediction_rankings]\n    tile_test_df['TopConfigs'] = idxs_string\n    tile_test_df = tile_test_df[['ID', 'TopConfigs']]\n    tile_test_df.head()\n    tile_test_df.to_csv(f'inference_{mode}_all.csv')\n\n    return tile_test_df","metadata":{"_uuid":"76c17ef6-c844-4919-b6af-9b3b65c929a1","_cell_guid":"4b66c5c0-6b72-4e54-9b34-32c164de3376","scrolled":true,"execution":{"iopub.status.busy":"2023-11-20T03:54:45.898221Z","iopub.execute_input":"2023-11-20T03:54:45.898757Z","iopub.status.idle":"2023-11-20T03:54:45.916543Z","shell.execute_reply.started":"2023-11-20T03:54:45.898689Z","shell.execute_reply":"2023-11-20T03:54:45.915474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tile_test_df = make_tile_prediction()","metadata":{"_uuid":"d462834d-632a-4d66-920d-e81615f3ec95","_cell_guid":"ae2e06dd-2dd2-4704-9bd0-e05bc0e7ac46","scrolled":true,"execution":{"iopub.status.busy":"2023-11-20T03:54:45.923799Z","iopub.execute_input":"2023-11-20T03:54:45.924191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"layout_test_df.head(3)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df = pd.read_csv('../input/predict-ai-model-runtime/sample_submission.csv')\n#tile_submission_df = submission_df.query(f\"ID not in {tile_test_df.ID.tolist()}\")\n#submission_df = pd.concat([tile_test_df, tile_submission_df])\n# tile_submission_df = submission_df.query(f\"ID not in {layout_test_df.ID.tolist()}\")\nsubmission_df = pd.concat([tile_test_df, layout_test_df])\nsubmission_df = submission_df[['ID', 'TopConfigs']]\nsubmission_df.to_csv('submission.csv', index=False)\nsubmission_df","metadata":{"_uuid":"d1690ded-823e-48a9-9e95-29b7cab147eb","_cell_guid":"1c2298eb-672e-4f2a-b981-aebae5815c10","scrolled":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # tile_df 5fold cv\n# from transformers import get_linear_schedule_with_warmup, get_cosine_schedule_with_warmup\n# print_step = 100\n# # get scheduler\n# def get_scheduler(params, optimizer, num_train_steps):\n#     if params['scheduler'] == 'linear':\n#         scheduler = get_linear_schedule_with_warmup(\n#             optimizer, num_training_steps=num_train_steps,\n#             num_warmup_steps=params['num_warmup_steps']\n#         )\n#     elif params['scheduler'] == 'cosine':\n#         scheduler = get_cosine_schedule_with_warmup(\n#             optimizer, num_training_steps=num_train_steps,\n#             num_warmup_steps=params['num_warmup_steps'], num_cycles=params['num_cycles']\n#         )\n#     return scheduler\n# # Training function\n# def train_fn(fold, train_loader, model, optimizer, epoch, scheduler, device, params):\n#     model.train()\n#     scaler = torch.cuda.amp.GradScaler(enabled=params['apex'])\n#     losses = []\n#     global_step = 0\n#     for step, batch in enumerate(train_loader):\n#         for k, v in batch.items():\n#             batch[k] = v.to(device)\n#         batch_size = len(batch.items())\n#         with torch.cuda.amp.autocast(enabled=params['apex']):\n#             output = model(**batch)\n# #         print(output)\n#         loss = output['loss']\n#         y_preds = output['outputs']\n#         y_order = output['order']\n#         if params['gradient_accumulation_steps'] > 1:\n#             loss = loss / params['gradient_accumulation_steps']\n#         losses.extend([loss.item()]* batch_size)\n#         scaler.scale(loss).backward()\n#         grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), params['max_grad_norm'])\n#         if (step + 1) % params['gradient_accumulation_steps'] == 0:\n#             scaler.step(optimizer)\n#             scaler.update()\n#             optimizer.zero_grad()\n#             global_step += 1\n#             if params['batch_scheduler']:\n#                 scheduler.step()\n#         # Log the message \n#         if step % print_step == 0 or step == (len(train_loader)-1):\n#             print(f'Epoch {epoch+1} [{step}/{len(train_loader)}] '\n#                   f'Loss: {loss.item():.4f}({np.mean(losses):.4f}) '\n#                   f'Grad: {grad_norm:.4f}  '\n#                   f'LR: {scheduler.get_lr()[0]:.8f}  ')\n#     return np.mean(losses)\n\n# # Validate the model and compute 'topk' metric\n# def valid_fn(valid_loader, model, topk, device, params):\n#     losses = []\n#     model.eval()\n#     preds = []\n#     for step, batch in enumerate(valid_loader):\n#         for k, v in batch.items():\n#             batch[k] = v.to(device)\n#         batch_size = len(batch.items())\n        \n#         with torch.no_grad():\n#             output = model(**batch)\n            \n#         loss = output['loss']\n#         y_preds = output['outputs']\n#         y_order = output['order']\n#         if params['gradient_accumulation_steps'] > 1:\n#             loss = loss / params['gradient_accumulation_steps']\n#         losses.extend([loss.item()]*batch_size)\n        \n#         config_attn_mask = torch.ones_like(batch['config_runtime'], device=batch['config_runtime'].device)\n#         topk.update(output['outputs'], batch['config_runtime'], config_attn_mask)\n#         # Log the message\n#         if step % print_step == 0 or step == (len(valid_loader)-1):\n#             print(f'EVAL: [{step}/{len(valid_loader)}] '\n#                    f'Loss: {loss.item():.4f}({np.mean(losses):.4f}) ')\n#     # cal metrics\n#     topk_res = topk.compute()\n#     print(f\"topk {topk_res:.3f}\")\n#     topk.reset()\n#     return np.mean(losses), topk_res","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # torch.set_float32_matmul_precision(\"medium\") # Set \n# device = 'cuda:0' if torch.cuda.is_available() else 'cpu'\n# OUTPUT_DIR = '/kaggle/working/'\n# config_kwargs = {'hidden_size': 128,\n#                  'num_attention_heads': 4,\n#                  'num_hidden_layers': 2,\n#                  'intermediate_size': 64,\n#                  'gradient_checkpointing': True,\n#                  'margin': 0.1,\n#                  'number_permutations': 4}\n# config = GraphConfig(**config_kwargs)\n\n# def training_objective(trial):\n#     params = {\"lr\": trial.suggest_float('lr', 1e-4, 1e-2),\n#               \"eps\": 1e-6,\n#               \"betas\": (0.9, 0.999),\n#               \"epochs\": 5*2,             # Epochs\n#               \"batch_size\": 8,\n#               \"scheduler\": 'linear', # cosine/linear\n#               \"num_cycles\": 0.5,\n#               \"num_warmup_steps\": 50,\n#               \"gradient_accumulation_steps\": 2,\n#               \"max_grad_norm\": 1000,\n#               \"seed\": 42,\n#               \"batch_scheduler\": True,\n#               \"apex\": False, # if True, loss may explode\n#               \"model\": \"BertGraph\"}\n#     # Create splits\n#     total_scores = []\n    \n#     # tile_train_df, tile_test_df = cv_split(tile_df, n_fold=params['n_fold'])\n#     for idx, fold in enumerate(range(n_fold)):\n#         # Create the model \n#         model = GraphEncoder(config)\n#         model.to(device)\n\n#         # Split the tile into train and valid datasets\n#         train_df = tile_train_df[tile_train_df['fold']==fold].reset_index(drop=True)\n#         valid_df = tile_train_df[tile_train_df['fold']!=fold].reset_index(drop=True)\n#         train_dataset = TileDataset(train_df, num_configs=24)\n#         valid_dataset = TileDataset(valid_df, num_configs=24)\n\n#         # Create dataloaders\n#         train_dataloader = DataLoader(train_dataset, collate_fn=collate_fn, batch_size=params['batch_size'], num_workers=2, shuffle=True, persistent_workers=True)\n#         valid_dataloader = DataLoader(valid_dataset, collate_fn=collate_fn, batch_size=params['batch_size'], num_workers=2)\n\n#         # loop\n#         optimizer_parameters = model.parameters()\n#         optimizer = AdamW(optimizer_parameters, lr=params['lr'], eps=params['eps'],\n#                           betas=params['betas'], maximize=False)\n\n#         num_train_steps = int(len(train_df) / params['batch_size'] * params['epochs'])\n#         scheduler = get_scheduler(params, optimizer, num_train_steps)\n\n#         topk = TileTopK()\n#         best_score = 1.0\n#         for epoch in range(params['epochs']):  # Find the best epochs\n#             start_time = time.time()\n#             # train\n#             avg_loss = train_fn(fold, train_dataloader, model, optimizer, epoch, scheduler, device, params)\n#             # eval\n#             avg_val_loss, topk_res = valid_fn(valid_dataloader, model, topk, device, params)\n\n#             # scoring\n#             # score = get_score(valid_labels, predictions)\n#             score = avg_val_loss\n#             elapsed = time.time() - start_time\n#             print(f'Fold {fold} Epoch {epoch+1} - avg_train_loss: {avg_loss:.4f}  avg_val_loss: {avg_val_loss:.4f}  time: {elapsed:.0f}s')\n            \n#             total_scores.append(score)\n#             if score < best_score:\n#                 best_score = score\n#                 print(f'Fold {fold} Epoch {epoch+1} - Save Best Score: {best_score:.4f} Model')\n#                 torch.save(model.state_dict(),\n#                            OUTPUT_DIR+f\"{params['model']}_fold{fold}_best.pth\")\n#         del model\n#         torch.cuda.empty_cache()\n#         gc.collect()\n        \n#     avg_score = np.mean(total_scores) # Use the average valid loss as objective\n#     print(f\"Average score {avg_score}\")\n#     return avg_score","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}