{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"\n\n##  Structure \n1. **Data Prep**: In this section, we load and preprocess the competition data.\n2. **Feature Engineering**: We generate and select relevant features for model training.\n3. **Model Training**: We train machine learning models on the prepared data.\n4. **Predict / Submit**: We make predictions on the test data and submit them for evaluation.\n","metadata":{}},{"cell_type":"code","source":"# 📚 Importing  libraries\nimport os  # For interacting with the operating system\nfrom pathlib import Path  # For working with file paths\nfrom typing import Dict, Optional, List, Union, Tuple  # For defining data types\nfrom dataclasses import dataclass  # For creating data classes\nimport math  # For mathematical operations\nimport numpy as np  # For numerical computations\nimport pandas as pd  # For data manipulation and analysis\nfrom datasets import Dataset  # For handling datasets\nfrom tqdm import tqdm  # For progress tracking\nimport torch  # For deep learning with PyTorch\nfrom torch import nn  # For neural network modules\nfrom torch.nn import functional as F  # For various functions used in neural networks\nfrom torch.nn.utils.rnn import pad_sequence  # For padding sequences\nfrom torch.utils.data import DataLoader  # For creating data loaders\nfrom transformers.modeling_outputs import BaseModelOutputWithPastAndCrossAttentions  # For transformer model outputs\nfrom transformers.pytorch_utils import apply_chunking_to_forward  # For handling chunking during model forward pass\nfrom transformers.activations import ACT2FN  # For transformer activations\nimport pytorch_lightning as pl  # For PyTorch Lightning, a useful library for training\nimport torchmetrics as tm  # For additional metrics\n\n\n# import bitsandbytes as bnb  \n","metadata":{"execution":{"iopub.status.busy":"2023-10-04T11:50:07.029359Z","iopub.execute_input":"2023-10-04T11:50:07.029744Z","iopub.status.idle":"2023-10-04T11:50:23.149952Z","shell.execute_reply.started":"2023-10-04T11:50:07.029719Z","shell.execute_reply":"2023-10-04T11:50:23.149064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## List all Files \nimport os\nos.listdir('/kaggle/input/predict-ai-model-runtime/npz_all/npz/layout/xla/random/train')\n","metadata":{"execution":{"iopub.status.busy":"2023-10-04T11:50:23.151935Z","iopub.execute_input":"2023-10-04T11:50:23.152864Z","iopub.status.idle":"2023-10-04T11:50:23.189743Z","shell.execute_reply.started":"2023-10-04T11:50:23.152825Z","shell.execute_reply":"2023-10-04T11:50:23.188757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define constants and configuration values","metadata":{}},{"cell_type":"code","source":"NODE_OP_CODES = 120  # Number of node operation codes\nNODE_FEATS = 140     # Number of node features\nCONFIG_FEATS = 24    # Number of configuration features\nNODE_CONFIG_FEATS = 18  # Number of combined node and configuration features\n","metadata":{"execution":{"iopub.status.busy":"2023-10-04T11:50:23.191087Z","iopub.execute_input":"2023-10-04T11:50:23.191664Z","iopub.status.idle":"2023-10-04T11:50:23.196549Z","shell.execute_reply.started":"2023-10-04T11:50:23.191633Z","shell.execute_reply":"2023-10-04T11:50:23.195291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":" ##  📄  Function to Generate Tile DataFrame","metadata":{}},{"cell_type":"code","source":"DATA_DIR = \"../input/predict-ai-model-runtime/npz_all/npz\"\n\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-10-04T11:50:23.199413Z","iopub.execute_input":"2023-10-04T11:50:23.200562Z","iopub.status.idle":"2023-10-04T11:50:23.208603Z","shell.execute_reply.started":"2023-10-04T11:50:23.200530Z","shell.execute_reply":"2023-10-04T11:50:23.207567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Generating and Displaying Tile DataFrame","metadata":{}},{"cell_type":"code","source":"# 🧩 Generate the tile DataFrame using the previously defined function\ntile_df = generate_tile_df()\n\n# Display the first few rows of the tile DataFrame\ntile_df.head(5)\n","metadata":{"execution":{"iopub.status.busy":"2023-10-04T11:50:23.210100Z","iopub.execute_input":"2023-10-04T11:50:23.210685Z","iopub.status.idle":"2023-10-04T11:50:36.750580Z","shell.execute_reply.started":"2023-10-04T11:50:23.210646Z","shell.execute_reply":"2023-10-04T11:50:36.749585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Definition of Functions and a Custom Dataset Class","metadata":{}},{"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 tile_loader(path):\n    tile_dict =  dict(np.load(path))\n    tile_dict = {k: torch.from_numpy(v) for k, v in tile_dict.items()}\n    tile_dict['edges_adjecency'] = edges_adjacency(tile_dict['edge_index'])\n    return tile_dict\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']])\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\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, 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 __getitem__(self, idx:int, selected_configs:List[int]=None):\n        tile_dict = 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","metadata":{"execution":{"iopub.status.busy":"2023-10-04T11:50:36.752288Z","iopub.execute_input":"2023-10-04T11:50:36.752640Z","iopub.status.idle":"2023-10-04T11:50:36.767695Z","shell.execute_reply.started":"2023-10-04T11:50:36.752605Z","shell.execute_reply":"2023-10-04T11:50:36.766670Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"\n\n","metadata":{}},{"cell_type":"code","source":"## Creating a Tile Dataset\ntile_dataset = TileDataset(tile_df)\n## Accessing Tile Data from the Dataset\nelem = tile_dataset[0]\nfor k,v in elem.items():\n    print(k, v.shape)","metadata":{"execution":{"iopub.status.busy":"2023-10-04T11:50:36.769328Z","iopub.execute_input":"2023-10-04T11:50:36.769984Z","iopub.status.idle":"2023-10-04T11:50:36.875540Z","shell.execute_reply.started":"2023-10-04T11:50:36.769952Z","shell.execute_reply":"2023-10-04T11:50:36.874608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"\n\n","metadata":{}},{"cell_type":"markdown","source":"## Data Preparation and Collation","metadata":{}},{"cell_type":"code","source":"def pad_edge_adjacency(edges_adjacency_list):\n    max_len = max([elem.shape[0] for elem in edges_adjacency_list])\n    return torch.stack([F.pad(elem, (0, max_len-elem.shape[0], 0, max_len-elem.shape[0]), value=0) for elem in edges_adjacency_list], dim=0)\n\n@dataclass\nclass LayoutCollator:\n    pad_to_multiple_of: int = 64\n    targets:bool = True\n    padding_idx:int = 120\n    node_padding_idx:int = 0\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(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":{"execution":{"iopub.status.busy":"2023-10-04T11:50:36.877023Z","iopub.execute_input":"2023-10-04T11:50:36.877585Z","iopub.status.idle":"2023-10-04T11:50:36.890140Z","shell.execute_reply.started":"2023-10-04T11:50:36.877551Z","shell.execute_reply":"2023-10-04T11:50:36.889155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"##  Collate batch  Function Using `LayoutCollator`","metadata":{}},{"cell_type":"code","source":"collate_fn = LayoutCollator(64)\nbatch = collate_fn([tile_dataset[0], tile_dataset[1]])\nfor k,v in batch.items():\n    print(k,v.shape)","metadata":{"execution":{"iopub.status.busy":"2023-10-04T11:50:36.891590Z","iopub.execute_input":"2023-10-04T11:50:36.891953Z","iopub.status.idle":"2023-10-04T11:50:36.932580Z","shell.execute_reply.started":"2023-10-04T11:50:36.891919Z","shell.execute_reply":"2023-10-04T11:50:36.931550Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## GraphConfig Dataclass","metadata":{}},{"cell_type":"code","source":"@dataclass\nclass GraphConfig:\n    num_hidden_layers: int = 8\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-12\n    hidden_act: str = 'gelu'\n    initializer_range: float = 0.02\n    output_hidden_states: bool = False\n    output_attentions: bool = False\n    gradient_checkpointing: bool = False\n    margin: float = 0.1\n    number_permutations: int = 10\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":{"execution":{"iopub.status.busy":"2023-10-04T11:50:36.936410Z","iopub.execute_input":"2023-10-04T11:50:36.937013Z","iopub.status.idle":"2023-10-04T11:50:36.946354Z","shell.execute_reply.started":"2023-10-04T11:50:36.936979Z","shell.execute_reply":"2023-10-04T11:50:36.945437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## `MultiElementRankLoss`","metadata":{}},{"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 generate_permutation(self,\n                             config_attn_mask: torch.Tensor\n                             ):\n        \"\"\"\n        Generate a permutation of the elements in the batch\n        Args:\n            config_attn_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            permutation: Tensor of shape (2, bs*seq_len) with the permutation of the elements\n        \"\"\"\n        num_elements = config_attn_mask.sum(1)\n        permutation_list = [torch.randperm(int(elem)) for elem in  num_elements.cpu().numpy()]\n        idxs_list = [num*torch.ones_like(elem) for num, elem in enumerate(permutation_list)]\n        permutation = torch.stack([torch.cat(idxs_list), torch.cat(permutation_list)])\n        return permutation\n    \n    def permute_tensor(self,\n                       tensor:torch.Tensor,\n                       permutation:torch.Tensor,\n                       x:torch.Tensor,\n                       y:torch.Tensor\n                       ):\n        \"\"\"\n        Permute the tensor according to the permutation\n        Args:\n            tensor: Tensor of shape (bs, seq_len) to be permuted\n            permutation: Tensor of shape (2, bs*seq_len) with the permutation of the elements\n            x: Tensor of shape (bs*seq_len) with the x coordinates of the elements to be permuted\n            y: Tensor of shape (bs*seq_len) with the y coordinates of the elements to be permuted\n        Returns:\n            permuted_tensor: Tensor of shape (bs, seq_len) with the permuted elements\n        \"\"\"\n        new_tensor = tensor.clone()\n        new_tensor[x, y] = new_tensor[permutation[0, :], permutation[1, :]]\n        return new_tensor\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        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,\n                outputs: torch.Tensor,\n                config_runtime: torch.Tensor,\n                config_idxs: torch.Tensor\n                ):\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":{"execution":{"iopub.status.busy":"2023-10-04T11:50:36.947781Z","iopub.execute_input":"2023-10-04T11:50:36.948385Z","iopub.status.idle":"2023-10-04T11:50:36.961203Z","shell.execute_reply.started":"2023-10-04T11:50:36.948353Z","shell.execute_reply":"2023-10-04T11:50:36.960258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## `TileTopK` Metric","metadata":{}},{"cell_type":"code","source":"\nclass TileTopK(tm.Metric):\n    \n    higher_is_better = True\n    \n    def __init__(self, k:int=5) -> None:\n        super().__init__()\n        self.add_state(\"runtimes\", default=[], dist_reduce_fx=None)\n        self.k = k\n        \n    def update(self, preds: torch.Tensor, target: torch.Tensor, config_attn_mask:torch.Tensor) -> None:\n        \"\"\"\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":{"execution":{"iopub.status.busy":"2023-10-04T11:50:36.962405Z","iopub.execute_input":"2023-10-04T11:50:36.963131Z","iopub.status.idle":"2023-10-04T11:50:36.976399Z","shell.execute_reply.started":"2023-10-04T11:50:36.963099Z","shell.execute_reply":"2023-10-04T11:50:36.975381Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Custom Model","metadata":{}},{"cell_type":"code","source":"# Modified from https://github.com/huggingface/transformers/blob/main/src/transformers/models/bert/modeling_bert.py\nclass 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(\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        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\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        mixed_query_layer = self.query(hidden_states)\n\n        # If this is instantiated as a cross-attention module, the keys\n        # and values come from an encoder; the attention mask needs to be\n        # such that the encoder's padding tokens are not attended to.\n\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.unsqueeze(0).repeat(self.config.num_hidden_layers, 1, 1, 1).unsqueeze(2),\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)\n    \n    \nclass 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, number_permutations=config.number_permutations)\n        \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":{"execution":{"iopub.status.busy":"2023-10-04T11:50:36.978151Z","iopub.execute_input":"2023-10-04T11:50:36.978821Z","iopub.status.idle":"2023-10-04T11:50:37.027613Z","shell.execute_reply.started":"2023-10-04T11:50:36.978788Z","shell.execute_reply":"2023-10-04T11:50:37.026330Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"##  `LightningWrapper` class","metadata":{}},{"cell_type":"code","source":"class LightningWrapper(pl.LightningModule):\n    def __init__(self, model:nn.Module):\n        super().__init__()\n        self.model = model\n        self.topk = TileTopK()\n        \n    def forward(self, x):\n        return self.model(x)\n\n    def training_step(self, batch, batch_idx):\n        outputs = self.model(**batch)\n        return outputs['loss']\n\n    def validation_step(self, batch, batch_idx):\n        outputs = self.model(**batch)\n        loss = outputs['loss']\n        self.log(\"val_loss\", loss, prog_bar=True)\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()\n        self.print(f\"topk {topk:.3f}\")\n        self.topk.reset()\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        optimizer = torch.optim.AdamW(self.trainer.model.parameters(), lr=1e-3)\n        return optimizer","metadata":{"execution":{"iopub.status.busy":"2023-10-04T11:50:37.029367Z","iopub.execute_input":"2023-10-04T11:50:37.029905Z","iopub.status.idle":"2023-10-04T11:50:37.044013Z","shell.execute_reply.started":"2023-10-04T11:50:37.029871Z","shell.execute_reply":"2023-10-04T11:50:37.043022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"config_kwargs = dict(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    )","metadata":{"execution":{"iopub.status.busy":"2023-10-04T11:50:37.045479Z","iopub.execute_input":"2023-10-04T11:50:37.046035Z","iopub.status.idle":"2023-10-04T11:50:37.062059Z","shell.execute_reply.started":"2023-10-04T11:50:37.046002Z","shell.execute_reply":"2023-10-04T11:50:37.061022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":" we use this config object to create an instance of your model with the desired architecture and hyperparameters.","metadata":{}},{"cell_type":"code","source":"config = GraphConfig(**config_kwargs)\nmodel = GraphEncoder(config)\nmodel = LightningWrapper(model)","metadata":{"execution":{"iopub.status.busy":"2023-10-04T11:50:37.063587Z","iopub.execute_input":"2023-10-04T11:50:37.064260Z","iopub.status.idle":"2023-10-04T11:50:37.082318Z","shell.execute_reply.started":"2023-10-04T11:50:37.064227Z","shell.execute_reply":"2023-10-04T11:50:37.081398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = tile_df.query(\"split == 'train'\").reset_index(drop=True)\nvalid_df = tile_df.query(\"split == 'valid'\").reset_index(drop=True)\ntrain_dataset = TileDataset(train_df, num_configs=24)\nvalid_dataset = TileDataset(valid_df, num_configs=24)","metadata":{"execution":{"iopub.status.busy":"2023-10-04T11:50:37.083758Z","iopub.execute_input":"2023-10-04T11:50:37.084338Z","iopub.status.idle":"2023-10-04T11:50:37.107718Z","shell.execute_reply.started":"2023-10-04T11:50:37.084307Z","shell.execute_reply":"2023-10-04T11:50:37.106814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataloader = DataLoader(train_dataset, collate_fn=collate_fn, batch_size=8, num_workers=2, shuffle=True, persistent_workers=True)\nvalid_dataloader = DataLoader(valid_dataset, collate_fn=collate_fn, batch_size=8, num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2023-10-04T11:50:37.109162Z","iopub.execute_input":"2023-10-04T11:50:37.109841Z","iopub.status.idle":"2023-10-04T11:50:37.115057Z","shell.execute_reply.started":"2023-10-04T11:50:37.109809Z","shell.execute_reply":"2023-10-04T11:50:37.114139Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## `trainer_config`","metadata":{}},{"cell_type":"code","source":"trainer_config = dict(\n    max_epochs= 50,\n    precision= 32,\n    gradient_clip_val= 1.0,\n    accumulate_grad_batches= 4,\n    check_val_every_n_epoch= 10)","metadata":{"execution":{"iopub.status.busy":"2023-10-04T11:50:37.116587Z","iopub.execute_input":"2023-10-04T11:50:37.117255Z","iopub.status.idle":"2023-10-04T11:50:37.129320Z","shell.execute_reply.started":"2023-10-04T11:50:37.117223Z","shell.execute_reply":"2023-10-04T11:50:37.128299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.set_float32_matmul_precision(\"medium\")\ntrainer = pl.Trainer(**trainer_config,)\ntrainer.fit(model, train_dataloader, valid_dataloader)","metadata":{"execution":{"iopub.status.busy":"2023-10-04T11:50:37.130993Z","iopub.execute_input":"2023-10-04T11:50:37.131658Z","iopub.status.idle":"2023-10-04T12:21:49.026734Z","shell.execute_reply.started":"2023-10-04T11:50:37.131622Z","shell.execute_reply":"2023-10-04T12:21:49.025101Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"split = 'test'\ntest_tile_df = tile_df.query(\"split == @split\").reset_index(drop=True)\ntest_tile_ds = TileDataset(test_tile_df, num_configs=-1)\ncollate_fn = LayoutCollator(64, targets=split!=\"test\")\ntest_dataloader = DataLoader(test_tile_ds, batch_size=1, shuffle=False, num_workers=0, collate_fn=collate_fn)","metadata":{"execution":{"iopub.status.busy":"2023-10-04T12:21:49.031513Z","iopub.execute_input":"2023-10-04T12:21:49.032544Z","iopub.status.idle":"2023-10-04T12:21:49.046804Z","shell.execute_reply.started":"2023-10-04T12:21:49.032508Z","shell.execute_reply":"2023-10-04T12:21:49.045940Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device='cuda:0' if torch.cuda.is_available() else 'cpu'\nmodel.to(device)\nmodel = model.eval()\n","metadata":{"execution":{"iopub.status.busy":"2023-10-04T12:21:49.048414Z","iopub.execute_input":"2023-10-04T12:21:49.052298Z","iopub.status.idle":"2023-10-04T12:21:52.248300Z","shell.execute_reply.started":"2023-10-04T12:21:49.052262Z","shell.execute_reply":"2023-10-04T12:21:52.247349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":" ## `chunk_batch` function","metadata":{}},{"cell_type":"code","source":"def chunk_batch(batch, start_idx, end_idx):\n    # Create an output dictionary to store the selected batch components.\n    output = {k: batch[k] for k in ['node_opcode', 'node_feat', 'edges_adjecency', 'node_attn_mask', 'node_config_ids']}\n    \n    # Slice the 'node_config_feat' component to create a smaller chunk.\n    output['node_config_feat'] = batch['node_config_feat'][:, start_idx: end_idx]\n    \n    return output\n","metadata":{"execution":{"iopub.status.busy":"2023-10-04T12:21:52.249994Z","iopub.execute_input":"2023-10-04T12:21:52.250708Z","iopub.status.idle":"2023-10-04T12:21:52.256957Z","shell.execute_reply.started":"2023-10-04T12:21:52.250673Z","shell.execute_reply":"2023-10-04T12:21:52.256065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_order = []\nfor batch in tqdm(test_dataloader):\n    batch.pop('config_idxs')\n    batch = {k: v.to(device) for k, v in batch.items()}\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.model(**chunked_batch)\n        chunk_order.extend(output['outputs'].cpu().numpy())\n    pred_order.append(np.argsort(np.concatenate(chunk_order))[:5])","metadata":{"execution":{"iopub.status.busy":"2023-10-04T12:21:52.258390Z","iopub.execute_input":"2023-10-04T12:21:52.258825Z","iopub.status.idle":"2023-10-04T12:23:30.567906Z","shell.execute_reply.started":"2023-10-04T12:21:52.258791Z","shell.execute_reply":"2023-10-04T12:23:30.566952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"idxs_string = [\";\".join(map(str,elem)) for elem in pred_order]\ntest_tile_df['TopConfigs'] = idxs_string\ntest_tile_df = test_tile_df[['ID', 'TopConfigs']]\ntest_tile_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-10-04T12:23:30.569223Z","iopub.execute_input":"2023-10-04T12:23:30.569582Z","iopub.status.idle":"2023-10-04T12:23:30.590173Z","shell.execute_reply.started":"2023-10-04T12:23:30.569548Z","shell.execute_reply":"2023-10-04T12:23:30.589063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## submission.csv","metadata":{}},{"cell_type":"code","source":"submission_df = pd.read_csv('../input/predict-ai-model-runtime/sample_submission.csv')\nsubmission_df = submission_df.query(f\"ID not in {test_tile_df.ID.tolist()}\")\nsubmission_df = pd.concat([test_tile_df, submission_df])\nsubmission_df.to_csv('submission.csv', index=False)\nsubmission_df","metadata":{"execution":{"iopub.status.busy":"2023-10-04T12:23:30.591473Z","iopub.execute_input":"2023-10-04T12:23:30.592786Z","iopub.status.idle":"2023-10-04T12:23:30.655018Z","shell.execute_reply.started":"2023-10-04T12:23:30.592749Z","shell.execute_reply":"2023-10-04T12:23:30.654119Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! cat submission.csv","metadata":{"execution":{"iopub.status.busy":"2023-10-04T12:29:56.335643Z","iopub.execute_input":"2023-10-04T12:29:56.336337Z","iopub.status.idle":"2023-10-04T12:29:57.498025Z","shell.execute_reply.started":"2023-10-04T12:29:56.336293Z","shell.execute_reply":"2023-10-04T12:29:57.496837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}