{"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":"code","source":"import os\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\nfrom tqdm import tqdm\nimport 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\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\nimport torchmetrics as tm\n# import bitsandbytes as bnb","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-09-15T23:01:44.287059Z","iopub.execute_input":"2023-09-15T23:01:44.287712Z","iopub.status.idle":"2023-09-15T23:01:59.877915Z","shell.execute_reply.started":"2023-09-15T23:01:44.287685Z","shell.execute_reply":"2023-09-15T23:01:59.876818Z"},"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-09-15T23:01:59.880601Z","iopub.execute_input":"2023-09-15T23:01:59.881297Z","iopub.status.idle":"2023-09-15T23:01:59.886448Z","shell.execute_reply.started":"2023-09-15T23:01:59.881256Z","shell.execute_reply":"2023-09-15T23:01:59.885346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-09-15T23:01:59.888085Z","iopub.execute_input":"2023-09-15T23:01:59.888725Z","iopub.status.idle":"2023-09-15T23:01:59.902853Z","shell.execute_reply.started":"2023-09-15T23:01:59.888687Z","shell.execute_reply":"2023-09-15T23:01:59.901918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tile_df = generate_tile_df()\ntile_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-09-15T23:01:59.907836Z","iopub.execute_input":"2023-09-15T23:01:59.908504Z","iopub.status.idle":"2023-09-15T23:02:10.684757Z","shell.execute_reply.started":"2023-09-15T23:01:59.908479Z","shell.execute_reply":"2023-09-15T23:02:10.683758Z"},"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 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']]) # 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, 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-09-15T23:02:10.686299Z","iopub.execute_input":"2023-09-15T23:02:10.686650Z","iopub.status.idle":"2023-09-15T23:02:10.706760Z","shell.execute_reply.started":"2023-09-15T23:02:10.686617Z","shell.execute_reply":"2023-09-15T23:02:10.704747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tile_dataset = TileDataset(tile_df)","metadata":{"execution":{"iopub.status.busy":"2023-09-15T23:02:10.709775Z","iopub.execute_input":"2023-09-15T23:02:10.710074Z","iopub.status.idle":"2023-09-15T23:02:10.732038Z","shell.execute_reply.started":"2023-09-15T23:02:10.710049Z","shell.execute_reply":"2023-09-15T23:02:10.731050Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"elem = tile_dataset[0]\nfor k,v in elem.items():\n    print(k, v.shape)","metadata":{"execution":{"iopub.status.busy":"2023-09-15T23:02:10.733225Z","iopub.execute_input":"2023-09-15T23:02:10.733615Z","iopub.status.idle":"2023-09-15T23:02:10.819822Z","shell.execute_reply.started":"2023-09-15T23:02:10.733582Z","shell.execute_reply":"2023-09-15T23:02:10.818784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"elem['edges_adjecency']\n","metadata":{"execution":{"iopub.status.busy":"2023-09-15T23:02:10.820909Z","iopub.execute_input":"2023-09-15T23:02:10.821448Z","iopub.status.idle":"2023-09-15T23:02:10.844418Z","shell.execute_reply.started":"2023-09-15T23:02:10.821415Z","shell.execute_reply":"2023-09-15T23:02:10.843558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-09-15T23:02:10.846798Z","iopub.execute_input":"2023-09-15T23:02:10.847372Z","iopub.status.idle":"2023-09-15T23:02:10.862942Z","shell.execute_reply.started":"2023-09-15T23:02:10.847339Z","shell.execute_reply":"2023-09-15T23:02:10.862071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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)\n","metadata":{"execution":{"iopub.status.busy":"2023-09-15T23:02:10.867451Z","iopub.execute_input":"2023-09-15T23:02:10.867707Z","iopub.status.idle":"2023-09-15T23:02:10.896529Z","shell.execute_reply.started":"2023-09-15T23:02:10.867685Z","shell.execute_reply":"2023-09-15T23:02:10.895666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-09-15T23:02:10.899622Z","iopub.execute_input":"2023-09-15T23:02:10.899904Z","iopub.status.idle":"2023-09-15T23:02:10.910928Z","shell.execute_reply.started":"2023-09-15T23:02:10.899858Z","shell.execute_reply":"2023-09-15T23:02:10.909889Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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,\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-09-15T23:02:10.912353Z","iopub.execute_input":"2023-09-15T23:02:10.912673Z","iopub.status.idle":"2023-09-15T23:02:10.924544Z","shell.execute_reply.started":"2023-09-15T23:02:10.912637Z","shell.execute_reply":"2023-09-15T23:02:10.923462Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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-09-15T23:02:10.925790Z","iopub.execute_input":"2023-09-15T23:02:10.926288Z","iopub.status.idle":"2023-09-15T23:02:10.940164Z","shell.execute_reply.started":"2023-09-15T23:02:10.926256Z","shell.execute_reply":"2023-09-15T23:02:10.939207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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        \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)\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\nclass 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-09-15T23:02:10.941695Z","iopub.execute_input":"2023-09-15T23:02:10.942219Z","iopub.status.idle":"2023-09-15T23:02:11.015626Z","shell.execute_reply.started":"2023-09-15T23:02:10.942188Z","shell.execute_reply":"2023-09-15T23:02:11.014716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-09-15T23:02:11.016939Z","iopub.execute_input":"2023-09-15T23:02:11.017464Z","iopub.status.idle":"2023-09-15T23:02:11.029642Z","shell.execute_reply.started":"2023-09-15T23:02:11.017431Z","shell.execute_reply":"2023-09-15T23:02:11.028565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config = GraphConfig(**config_kwargs)","metadata":{"execution":{"iopub.status.busy":"2023-09-15T23:02:11.030980Z","iopub.execute_input":"2023-09-15T23:02:11.031322Z","iopub.status.idle":"2023-09-15T23:02:11.043429Z","shell.execute_reply.started":"2023-09-15T23:02:11.031291Z","shell.execute_reply":"2023-09-15T23:02:11.042451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = GraphEncoder(config)\nmodel = LightningWrapper(model)","metadata":{"execution":{"iopub.status.busy":"2023-09-15T23:02:11.044549Z","iopub.execute_input":"2023-09-15T23:02:11.044787Z","iopub.status.idle":"2023-09-15T23:02:11.061955Z","shell.execute_reply.started":"2023-09-15T23:02:11.044766Z","shell.execute_reply":"2023-09-15T23:02:11.060912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tile_df\n","metadata":{"execution":{"iopub.status.busy":"2023-09-15T23:02:11.063244Z","iopub.execute_input":"2023-09-15T23:02:11.063567Z","iopub.status.idle":"2023-09-15T23:02:11.082320Z","shell.execute_reply.started":"2023-09-15T23:02:11.063537Z","shell.execute_reply":"2023-09-15T23:02:11.081376Z"},"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-09-15T23:02:11.084136Z","iopub.execute_input":"2023-09-15T23:02:11.084969Z","iopub.status.idle":"2023-09-15T23:02:11.104825Z","shell.execute_reply.started":"2023-09-15T23:02:11.084795Z","shell.execute_reply":"2023-09-15T23:02:11.103950Z"},"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-09-15T23:02:11.105935Z","iopub.execute_input":"2023-09-15T23:02:11.106845Z","iopub.status.idle":"2023-09-15T23:02:11.112474Z","shell.execute_reply.started":"2023-09-15T23:02:11.106812Z","shell.execute_reply":"2023-09-15T23:02:11.111340Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-09-15T23:02:11.113964Z","iopub.execute_input":"2023-09-15T23:02:11.114317Z","iopub.status.idle":"2023-09-15T23:02:11.123983Z","shell.execute_reply.started":"2023-09-15T23:02:11.114286Z","shell.execute_reply":"2023-09-15T23:02:11.123138Z"},"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-09-15T23:02:11.125357Z","iopub.execute_input":"2023-09-15T23:02:11.125711Z","iopub.status.idle":"2023-09-15T23:30:07.548929Z","shell.execute_reply.started":"2023-09-15T23:02:11.125652Z","shell.execute_reply":"2023-09-15T23:30:07.547951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda:0' if torch.cuda.is_available() else 'cpu'\nsplit = '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)\nmodel.to(device)\nmodel = model.eval()\ndef chunk_batch(batch, start_idx, end_idx):\n    output = {k:batch[k] for k in ['node_opcode', 'node_feat', 'edges_adjecency', 'node_attn_mask', 'node_config_ids']}\n    output['node_config_feat'] = batch['node_config_feat'][:, start_idx: end_idx]\n    return output\n    \npred_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    # Chunk the configs to avoid OOM errors\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-09-15T23:30:07.550911Z","iopub.execute_input":"2023-09-15T23:30:07.551251Z","iopub.status.idle":"2023-09-15T23:31:24.586398Z","shell.execute_reply.started":"2023-09-15T23:30:07.551215Z","shell.execute_reply":"2023-09-15T23:31:24.584769Z"},"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-09-15T23:31:24.587631Z","iopub.execute_input":"2023-09-15T23:31:24.588272Z","iopub.status.idle":"2023-09-15T23:31:24.609423Z","shell.execute_reply.started":"2023-09-15T23:31:24.588235Z","shell.execute_reply":"2023-09-15T23:31:24.608323Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-09-15T23:31:24.610570Z","iopub.execute_input":"2023-09-15T23:31:24.611227Z","iopub.status.idle":"2023-09-15T23:31:24.678735Z","shell.execute_reply.started":"2023-09-15T23:31:24.611189Z","shell.execute_reply":"2023-09-15T23:31:24.677656Z"},"trusted":true},"execution_count":null,"outputs":[]}]}